Merge pull request #472 from begriffs/hasql-19

Upgrade to Hasql 19
This commit is contained in:
Joe Nelson
2016-01-27 14:59:08 -08:00
23 changed files with 745 additions and 669 deletions
+1 -1
View File
@@ -8,7 +8,7 @@ dependencies:
override: override:
- cabal update - cabal update
- cabal sandbox init - cabal sandbox init
- cabal install --upgrade-dependencies --constraint="template-haskell installed" --dependencies-only --enable-tests - cabal install --upgrade-dependencies --constraint="template-haskell installed" --dependencies-only --enable-tests --reorder-goals
- cabal configure --enable-tests -f ci - cabal configure --enable-tests -f ci
test: test:
post: post:
+12 -10
View File
@@ -28,7 +28,7 @@ executable postgrest
ghc-options: -Wall -W -O2 ghc-options: -Wall -W -O2
main-is: PostgREST/Main.hs main-is: PostgREST/Main.hs
default-extensions: OverloadedStrings, ScopedTypeVariables, QuasiQuotes default-extensions: OverloadedStrings, ScopedTypeVariables, QuasiQuotes, LambdaCase
default-language: Haskell2010 default-language: Haskell2010
build-depends: aeson >= 0.8 build-depends: aeson >= 0.8
, base >= 4.8 && < 5 , base >= 4.8 && < 5
@@ -36,15 +36,16 @@ executable postgrest
, case-insensitive , case-insensitive
, cassava , cassava
, containers , containers
, contravariant
, errors , errors
, hasql >= 0.7.3 && < 0.8 , hasql >= 0.19.3.3 && < 0.20
, hasql-backend >= 0.4.1 && < 0.5 , interpolatedstring-perl6
, hasql-postgres >= 0.10.4 && < 0.11
, jwt , jwt
, 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
@@ -57,7 +58,7 @@ executable postgrest
, wai-cors , wai-cors
, wai-extra , wai-extra
, wai-middleware-static >= 0.6.0 , wai-middleware-static >= 0.6.0
, warp >= 3.0.2 , warp >= 3.1.0
, HTTP, http-types , HTTP, http-types
, MissingH , MissingH
, Ranged-sets , Ranged-sets
@@ -92,11 +93,11 @@ library
, case-insensitive , case-insensitive
, cassava , cassava
, containers , containers
, contravariant
, errors , errors
, hasql , hasql
, hasql-backend
, hasql-postgres
, http-types , http-types
, interpolatedstring-perl6
, jwt , jwt
, optparse-applicative , optparse-applicative
, parsec , parsec
@@ -133,7 +134,7 @@ 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 default-extensions: OverloadedStrings, ScopedTypeVariables, QuasiQuotes, LambdaCase
Hs-Source-Dirs: test, src Hs-Source-Dirs: test, src
if flag(ci) if flag(ci)
ghc-options: -Wall -W -Werror ghc-options: -Wall -W -Werror
@@ -168,22 +169,23 @@ Test-Suite spec
, case-insensitive , case-insensitive
, cassava , cassava
, containers , containers
, contravariant
, errors , errors
, hasql , hasql
, hasql-backend
, hasql-postgres
, heredoc , heredoc
, hlint , hlint
, hspec == 2.2.* , hspec == 2.2.*
, hspec-wai , hspec-wai
, hspec-wai-json , hspec-wai-json
, http-types , http-types
, interpolatedstring-perl6
, jwt , jwt
, optparse-applicative , optparse-applicative
, packdeps , packdeps
, parsec , parsec
, process , process
, regex-tdfa , regex-tdfa
, resource-pool
, safe , safe
, scientific , scientific
, string-conversions , string-conversions
+34 -40
View File
@@ -10,8 +10,6 @@ import Control.Applicative
import Control.Arrow ((***)) import Control.Arrow ((***))
import Control.Monad (join) import Control.Monad (join)
import Data.Bifunctor (first) import Data.Bifunctor (first)
import qualified Data.ByteString.Lazy as BL
import Data.Functor.Identity
import Data.List (find, sortBy, delete) import Data.List (find, sortBy, delete)
import Data.Maybe (fromMaybe, fromJust, mapMaybe) import Data.Maybe (fromMaybe, fromJust, mapMaybe)
import Data.Ord (comparing) import Data.Ord (comparing)
@@ -33,9 +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 as H import qualified Hasql.Session as H
import qualified Hasql.Backend as B
import qualified Hasql.Postgres as P
import PostgREST.Config (AppConfig (..)) import PostgREST.Config (AppConfig (..))
import PostgREST.Parsers import PostgREST.Parsers
@@ -49,8 +45,7 @@ import PostgREST.Types
import PostgREST.Auth (tokenJWT) import PostgREST.Auth (tokenJWT)
import PostgREST.Error (errResponse) import PostgREST.Error (errResponse)
import PostgREST.QueryBuilder ( asJson import PostgREST.QueryBuilder ( callProc
, callProc
, addJoinConditions , addJoinConditions
, sourceCTEName , sourceCTEName
, requestToQuery , requestToQuery
@@ -58,11 +53,12 @@ import PostgREST.QueryBuilder ( asJson
, addRelations , addRelations
, createReadStatement , createReadStatement
, createWriteStatement , createWriteStatement
, ResultsWithCount
) )
import Prelude import Prelude
app :: DbStructure -> AppConfig -> RequestBody -> Request -> H.Tx P.Postgres s Response app :: DbStructure -> AppConfig -> RequestBody -> Request -> H.Session 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)
@@ -82,17 +78,17 @@ app dbStructure conf reqBody req =
if range == emptyRange if range == emptyRange
then return $ errResponse status416 "HTTP Range error" then return $ errResponse status416 "HTTP Range error"
else do else do
row <- H.maybeEx stm row <- H.query () stm
let (tableTotal, queryTotal, _ , body) = extractQueryResult row let (tableTotal, queryTotal, _ , body) = row
if singular if singular
then return $ if queryTotal <= 0 then return $ if queryTotal <= 0
then responseLBS status404 [] "" then responseLBS status404 [] ""
else responseLBS status200 [contentTypeH] (fromMaybe "{}" body) else responseLBS status200 [contentTypeH] (cs body)
else do else do
let frm = rangeOffset range let frm = toInteger $ rangeOffset range
to = frm+queryTotal-1 to = frm + toInteger queryTotal - 1
contentRange = contentRangeH frm to tableTotal contentRange = contentRangeH frm to (toInteger <$> tableTotal)
status = rangeStatus frm to tableTotal status = rangeStatus frm to (toInteger <$> tableTotal)
canonical = urlEncodeVars -- should this be moved to the dbStructure (location)? canonical = urlEncodeVars -- should this be moved to the dbStructure (location)?
. sortBy (comparing fst) . sortBy (comparing fst)
. map (join (***) cs) . map (join (***) cs)
@@ -104,46 +100,47 @@ app dbStructure conf reqBody req =
"/" <> cs (qiName qi) <> "/" <> cs (qiName qi) <>
if Prelude.null canonical then "" else "?" <> cs canonical if Prelude.null canonical then "" else "?" <> cs canonical
) )
] (fromMaybe "[]" body) ] (cs body)
(ActionCreate, TargetIdent qi@(QualifiedIdentifier _ table), (ActionCreate, TargetIdent qi@(QualifiedIdentifier _ table),
Just payload@(PayloadJSON (UniformObjects rows))) -> Just payload@(PayloadJSON uniform@(UniformObjects rows))) ->
case mutateSqlParts of case mutateSqlParts of
Left e -> return $ responseLBS status400 [jsonH] $ cs e Left e -> return $ responseLBS status400 [jsonH] $ cs e
Right (sq,mq) -> do Right (sq,mq) -> do
let isSingle = (==1) $ V.length rows let isSingle = (==1) $ V.length rows
let pKeys = map pkName $ filter (filterPk schema table) allPrKeys -- would it be ok to move primary key detection in the query itself? let pKeys = map pkName $ filter (filterPk schema table) allPrKeys -- would it be ok to move primary key detection in the query itself?
let stm = createWriteStatement qi sq mq isSingle (iPreferRepresentation apiRequest) pKeys (contentType == TextCSV) payload let stm = createWriteStatement qi sq mq isSingle (iPreferRepresentation apiRequest) pKeys (contentType == TextCSV) payload
row <- H.maybeEx stm row <- H.query uniform stm
let (_, _, location, body) = extractQueryResult row let (_, _, location, body) = extractQueryResult row
return $ responseLBS status201 return $ responseLBS status201
[ [
contentTypeH, contentTypeH,
(hLocation, "/" <> cs table <> "?" <> cs (fromMaybe "" location)) (hLocation, "/" <> cs table <> "?" <> cs location)
] ]
$ if iPreferRepresentation apiRequest == Full then fromMaybe "[]" body else "" $ if iPreferRepresentation apiRequest == Full then cs body else ""
(ActionUpdate, TargetIdent qi, Just payload@(PayloadJSON _)) -> (ActionUpdate, TargetIdent qi, Just payload@(PayloadJSON uniform)) ->
case mutateSqlParts of case mutateSqlParts of
Left e -> return $ responseLBS status400 [jsonH] $ cs e Left e -> return $ responseLBS status400 [jsonH] $ cs e
Right (sq,mq) -> do Right (sq,mq) -> do
let stm = createWriteStatement qi sq mq False (iPreferRepresentation apiRequest) [] (contentType == TextCSV) payload let stm = createWriteStatement qi sq mq False (iPreferRepresentation apiRequest) [] (contentType == TextCSV) payload
row <- H.maybeEx stm row <- H.query uniform stm
let (_, queryTotal, _, body) = extractQueryResult row let (_, queryTotal, _, body) = extractQueryResult row
r = contentRangeH 0 (queryTotal-1) (Just queryTotal) r = contentRangeH 0 (toInteger $ queryTotal-1) (toInteger <$> Just queryTotal)
s = case () of _ | queryTotal == 0 -> status404 s = case () of _ | queryTotal == 0 -> status404
| iPreferRepresentation apiRequest == Full -> status200 | iPreferRepresentation apiRequest == Full -> status200
| otherwise -> status204 | otherwise -> status204
return $ responseLBS s [contentTypeH, r] return $ responseLBS s [contentTypeH, r]
$ if iPreferRepresentation apiRequest == Full then fromMaybe "[]" body else "" $ if iPreferRepresentation apiRequest == Full then cs body else ""
(ActionDelete, TargetIdent qi, Nothing) -> (ActionDelete, TargetIdent qi, Nothing) ->
case mutateSqlParts of case mutateSqlParts of
Left e -> return $ responseLBS status400 [jsonH] $ cs e Left e -> return $ responseLBS status400 [jsonH] $ cs e
Right (sq,mq) -> do Right (sq,mq) -> do
let fakeload = PayloadJSON $ UniformObjects V.empty let emptyUniform = UniformObjects V.empty
let fakeload = PayloadJSON emptyUniform
let stm = createWriteStatement qi sq mq False (iPreferRepresentation apiRequest) [] (contentType == TextCSV) fakeload let stm = createWriteStatement qi sq mq False (iPreferRepresentation apiRequest) [] (contentType == TextCSV) fakeload
row <- H.maybeEx stm row <- H.query emptyUniform stm
let (_, queryTotal, _, _) = extractQueryResult row let (_, queryTotal, _, _) = extractQueryResult row
return $ if queryTotal == 0 return $ if queryTotal == 0
then notFound then notFound
@@ -154,31 +151,29 @@ app dbStructure conf reqBody req =
pkeys = map pkName $ filter (filterPk tSchema tTable) allPrKeys pkeys = map pkName $ filter (filterPk tSchema tTable) allPrKeys
body = encode (TableOptions cols pkeys) body = encode (TableOptions cols pkeys)
filterCol :: Schema -> TableName -> Column -> Bool filterCol :: Schema -> TableName -> Column -> Bool
filterCol sc tb (Column{colTable=Table{tableSchema=s, tableName=t}}) = s==sc && t==tb filterCol sc tb Column{colTable=Table{tableSchema=s, tableName=t}} = s==sc && t==tb
filterCol _ _ _ = False filterCol _ _ _ = False
return $ responseLBS status200 [jsonH, allOrigins] $ cs body return $ responseLBS status200 [jsonH, allOrigins] $ cs body
(ActionInvoke, TargetIdent qi, (ActionInvoke, TargetIdent qi,
Just (PayloadJSON (UniformObjects payload))) -> do Just (PayloadJSON (UniformObjects payload))) -> do
exists <- doesProcExist qi exists <- H.query qi doesProcExist
if exists if exists
then do then do
let p = V.head payload let p = V.head payload
call = B.Stmt "select " V.empty True <>
asJson (callProc qi p)
jwtSecret = configJwtSecret conf jwtSecret = configJwtSecret conf
bodyJson :: Maybe (Identity Value) <- H.maybeEx call bodyJson <- H.query () (callProc qi p)
returnJWT <- doesProcReturnJWT qi returnJWT <- H.query qi doesProcReturnJWT
return $ responseLBS status200 [jsonH] return $ responseLBS status200 [jsonH]
(let body = fromMaybe emptyArray $ runIdentity <$> bodyJson in (let body = fromMaybe emptyArray bodyJson in
if returnJWT if returnJWT
then "{\"token\":\"" <> cs (tokenJWT jwtSecret body) <> "\"}" then "{\"token\":\"" <> cs (tokenJWT jwtSecret body) <> "\"}"
else cs $ encode body) else cs $ encode body)
else return notFound else return notFound
(ActionRead, TargetRoot, Nothing) -> do (ActionRead, TargetRoot, Nothing) -> do
body <- encode <$> accessibleTables (cs schema) body <- encode <$> H.query schema accessibleTables
return $ responseLBS status200 [jsonH] $ cs body return $ responseLBS status200 [jsonH] $ cs body
(ActionUnknown _, _, _) -> return notFound (ActionUnknown _, _, _) -> return notFound
@@ -206,14 +201,14 @@ app dbStructure conf reqBody req =
readSqlParts = (,) <$> selectQuery <*> countQuery readSqlParts = (,) <$> selectQuery <*> countQuery
mutateSqlParts = (,) <$> selectQuery <*> mutateQuery mutateSqlParts = (,) <$> selectQuery <*> mutateQuery
rangeStatus :: Int -> Int -> Maybe Int -> Status rangeStatus :: Integer -> Integer -> Maybe Integer -> Status
rangeStatus _ _ Nothing = status200 rangeStatus _ _ Nothing = status200
rangeStatus frm to (Just total) rangeStatus frm to (Just total)
| frm > total = status416 | frm > total = status416
| (1 + to - frm) < total = status206 | (1 + to - frm) < total = status206
| otherwise = status200 | otherwise = status200
contentRangeH :: Int -> Int -> Maybe Int -> Header contentRangeH :: Integer -> Integer -> Maybe Integer -> Header
contentRangeH frm to total = contentRangeH frm to total =
("Content-Range", cs headerValue) ("Content-Range", cs headerValue)
where where
@@ -298,7 +293,7 @@ buildMutateRequest apiRequest =
cond = first formatParserError $ map snd <$> mapM pRequestFilter mutateFilters cond = first formatParserError $ map snd <$> mapM pRequestFilter mutateFilters
addFilter :: (Path, Filter) -> ReadRequest -> ReadRequest addFilter :: (Path, Filter) -> ReadRequest -> ReadRequest
addFilter ([], flt) (Node (q@(Select {flt_=flts}), i) forest) = Node (q {flt_=flt:flts}, i) forest addFilter ([], flt) (Node (q@Select {flt_=flts}, i) forest) = Node (q {flt_=flt:flts}, i) forest
addFilter (path, flt) (Node rn forest) = addFilter (path, flt) (Node rn forest) =
case targetNode of case targetNode of
Nothing -> Node rn forest -- the filter is silenty dropped in the Request does not contain the required path Nothing -> Node rn forest -- the filter is silenty dropped in the Request does not contain the required path
@@ -335,6 +330,5 @@ instance ToJSON TableOptions where
, "pkey" .= tblOptpkey t ] , "pkey" .= tblOptpkey t ]
extractQueryResult :: Maybe (Maybe Int, Int, Maybe BL.ByteString, Maybe BL.ByteString) extractQueryResult :: Maybe ResultsWithCount -> ResultsWithCount
-> (Maybe Int, Int, Maybe BL.ByteString, Maybe BL.ByteString) extractQueryResult = fromMaybe (Nothing, 0, "", "")
extractQueryResult = fromMaybe (Just 0, 0, Just "", Just "")
+7 -6
View File
@@ -21,6 +21,7 @@ module PostgREST.Auth (
import Control.Monad (join) import Control.Monad (join)
import Data.Aeson (Value (..), Object) import Data.Aeson (Value (..), Object)
import Data.Aeson.Types (emptyObject, emptyArray) import Data.Aeson.Types (emptyObject, emptyArray)
import qualified Data.ByteString as BS
import Data.Vector as V (null, head) import Data.Vector as V (null, head)
import Data.Map as M (fromList, toList) import Data.Map as M (fromList, toList)
import Data.Monoid ((<>)) import Data.Monoid ((<>))
@@ -38,12 +39,12 @@ import qualified Data.HashMap.Lazy as H
this one is mapped to a SET ROLE statement. this one is mapped to a SET ROLE statement.
In case there is any problem decoding the JWT it returns Nothing. In case there is any problem decoding the JWT it returns Nothing.
-} -}
claimsToSQL :: JWT.ClaimsMap -> [Text] claimsToSQL :: JWT.ClaimsMap -> [BS.ByteString]
claimsToSQL = map setVar . toList claimsToSQL = map setVar . toList
where where
setVar ("role", String val) = setRole val setVar ("role", String val) = setRole val
setVar (k, val) = "set local postgrest.claims." <> pgFmtIdent k <> setVar (k, val) = "set local postgrest.claims." <> cs (pgFmtIdent k) <>
" = " <> valueToVariable val <> ";" " = " <> cs (valueToVariable val) <> ";"
valueToVariable = pgFmtLit . unquoted valueToVariable = pgFmtLit . unquoted
{-| {-|
@@ -65,9 +66,9 @@ jwtClaims secret input time =
claim prop = prop . JWT.claims <$> decoded claim prop = prop . JWT.claims <$> decoded
customClaims = claim JWT.unregisteredClaims customClaims = claim JWT.unregisteredClaims
-- | Receives the name of a role and returns a SET ROLE statement {-| Receives the name of a role and returns a SET ROLE statement -}
setRole :: Text -> Text setRole :: Text -> BS.ByteString
setRole role = "set local role " <> cs (pgFmtLit role) <> ";" setRole r = "set local role " <> cs (pgFmtLit r) <> ";"
{-| {-|
+1 -1
View File
@@ -42,7 +42,7 @@ data AppConfig = AppConfig {
, configSchema :: String , configSchema :: String
, configJwtSecret :: Secret , configJwtSecret :: Secret
, configPool :: Int , configPool :: Int
, configMaxRows :: Maybe Int , configMaxRows :: Maybe Integer
} }
argParser :: Parser AppConfig argParser :: Parser AppConfig
+391 -337
View File
@@ -10,28 +10,32 @@ module PostgREST.DbStructure (
, doesProcReturnJWT , doesProcReturnJWT
) where ) where
import qualified Hasql.Query as H
import qualified Hasql.Encoders as HE
import qualified Hasql.Decoders as HD
import Control.Applicative import Control.Applicative
import Control.Monad (join) import Control.Monad (join, replicateM)
import Data.Functor.Identity import Data.Functor.Contravariant (contramap)
import Text.InterpolatedString.Perl6 (q)
import Data.List (elemIndex, find, subsequences, sort, transpose) import Data.List (elemIndex, find, subsequences, sort, transpose)
import Data.Maybe (fromMaybe, fromJust, isJust, mapMaybe, listToMaybe) import Data.Maybe (fromMaybe, fromJust, isJust, mapMaybe, listToMaybe)
import Data.Monoid import Data.Monoid
import Data.Text (Text, split) import Data.Text (Text, split)
import qualified Hasql as H import qualified Hasql.Session as H
import qualified Hasql.Postgres as P
import qualified Hasql.Backend as B
import PostgREST.Types import PostgREST.Types
import GHC.Exts (groupWith) import GHC.Exts (groupWith)
import Data.Int (Int32)
import Prelude import Prelude
getDbStructure :: Schema -> H.Tx P.Postgres s DbStructure getDbStructure :: Schema -> H.Session DbStructure
getDbStructure schema = do getDbStructure schema = do
tabs <- allTables tabs <- H.query () allTables
cols <- allColumns tabs cols <- H.query () $ allColumns tabs
syns <- allSynonyms cols syns <- H.query () $ allSynonyms cols
rels <- allRelations tabs cols rels <- H.query () $ allRelations tabs cols
keys <- allPrimaryKeys tabs keys <- H.query () $ allPrimaryKeys tabs
let rels' = (addManyToManyRelations . raiseRelations schema syns . addParentRelations . addSynonymousRelations syns) rels let rels' = (addManyToManyRelations . raiseRelations schema syns . addParentRelations . addSynonymousRelations syns) rels
cols' = addForeignKeys rels' cols cols' = addForeignKeys rels' cols
@@ -44,60 +48,113 @@ getDbStructure schema = do
, dbPrimaryKeys = keys' , dbPrimaryKeys = keys'
} }
doesProc :: forall c s. B.CxValue c Int => encodeQi :: HE.Params QualifiedIdentifier
(Text -> Text -> B.Stmt c) -> QualifiedIdentifier -> H.Tx c s Bool encodeQi =
doesProc stmt qi = do contramap qiSchema (HE.value HE.text) <>
row :: Maybe (Identity Int) <- H.maybeEx $ stmt (qiSchema qi) (qiName qi) contramap qiName (HE.value HE.text)
return $ isJust row
doesProcExist :: QualifiedIdentifier -> H.Tx P.Postgres s Bool decodeTables :: HD.Result [Table]
doesProcExist = doesProc [H.stmt| decodeTables =
HD.rowsList tblRow
where
tblRow = Table <$> HD.value HD.text <*> HD.value HD.text
<*> HD.value HD.bool
decodeColumns :: [Table] -> HD.Result [Column]
decodeColumns tables =
mapMaybe (columnFromRow tables) <$> HD.rowsList colRow
where
colRow =
(,,,,,,,,,,)
<$> HD.value HD.text <*> HD.value HD.text
<*> HD.value HD.text <*> HD.value HD.int4
<*> HD.value HD.bool <*> HD.value HD.text
<*> HD.value HD.bool
<*> HD.nullableValue HD.int4
<*> HD.nullableValue HD.int4
<*> HD.nullableValue HD.text
<*> HD.nullableValue HD.text
decodeRelations :: [Table] -> [Column] -> HD.Result [Relation]
decodeRelations tables cols =
mapMaybe (relationFromRow tables cols) <$> HD.rowsList relRow
where
relRow = (,,,,,)
<$> HD.value HD.text
<*> HD.value HD.text
<*> HD.value (HD.array (HD.arrayDimension replicateM (HD.arrayValue HD.text)))
<*> HD.value HD.text
<*> HD.value HD.text
<*> HD.value (HD.array (HD.arrayDimension replicateM (HD.arrayValue HD.text)))
decodePks :: [Table] -> HD.Result [PrimaryKey]
decodePks tables =
mapMaybe (pkFromRow tables) <$> HD.rowsList pkRow
where
pkRow = (,,) <$> HD.value HD.text <*> HD.value HD.text <*> HD.value HD.text
decodeSynonyms :: [Column] -> HD.Result [(Column,Column)]
decodeSynonyms cols =
mapMaybe (synonymFromRow cols) <$> HD.rowsList synRow
where
synRow = (,,,,,)
<$> HD.value HD.text <*> HD.value HD.text
<*> HD.value HD.text <*> HD.value HD.text
<*> HD.value HD.text <*> HD.value HD.text
doesProcExist :: H.Query QualifiedIdentifier Bool
doesProcExist =
H.statement sql encodeQi (HD.singleRow (HD.value HD.bool)) True
where
sql = [q| SELECT EXISTS (
SELECT 1 SELECT 1
FROM pg_catalog.pg_namespace n FROM pg_catalog.pg_namespace n
JOIN pg_catalog.pg_proc p JOIN pg_catalog.pg_proc p
ON pronamespace = n.oid ON pronamespace = n.oid
WHERE nspname = ? WHERE nspname = $1
AND proname = ? AND proname = $2
|] ) |]
doesProcReturnJWT :: QualifiedIdentifier -> H.Tx P.Postgres s Bool doesProcReturnJWT :: H.Query QualifiedIdentifier Bool
doesProcReturnJWT = doesProc [H.stmt| doesProcReturnJWT =
H.statement sql encodeQi (HD.singleRow (HD.value HD.bool)) True
where
sql = [q| SELECT EXISTS (
SELECT 1 SELECT 1
FROM pg_catalog.pg_namespace n FROM pg_catalog.pg_namespace n
JOIN pg_catalog.pg_proc p JOIN pg_catalog.pg_proc p
ON pronamespace = n.oid ON pronamespace = n.oid
WHERE nspname = ? WHERE nspname = $1
AND proname = ? AND proname = $2
AND pg_catalog.pg_get_function_result(p.oid) like '%jwt_claims' AND pg_catalog.pg_get_function_result(p.oid) like '%jwt_claims'
|] ) |]
accessibleTables :: Schema -> H.Tx P.Postgres s [Table] accessibleTables :: H.Query Schema [Table]
accessibleTables schema = do accessibleTables =
rows <- H.listEx $ H.statement sql (HE.value HE.text) decodeTables True
[H.stmt| where
select sql = [q|
n.nspname as table_schema, select
relname as table_name, n.nspname as table_schema,
c.relkind = 'r' or (c.relkind IN ('v', 'f')) and (pg_relation_is_updatable(c.oid::regclass, false) & 8) = 8 relname as table_name,
or (exists ( c.relkind = 'r' or (c.relkind IN ('v', 'f')) and (pg_relation_is_updatable(c.oid::regclass, false) & 8) = 8
select 1 or (exists (
from pg_trigger select 1
where pg_trigger.tgrelid = c.oid and (pg_trigger.tgtype::integer & 69) = 69) from pg_trigger
) as insertable where pg_trigger.tgrelid = c.oid and (pg_trigger.tgtype::integer & 69) = 69)
from ) as insertable
pg_class c from
join pg_namespace n on n.oid = c.relnamespace pg_class c
where join pg_namespace n on n.oid = c.relnamespace
c.relkind in ('v', 'r', 'm') where
and n.nspname = ? c.relkind in ('v', 'r', 'm')
and ( and n.nspname = $1
pg_has_role(c.relowner, 'USAGE'::text) and (
or has_table_privilege(c.oid, 'SELECT, INSERT, UPDATE, DELETE, TRUNCATE, REFERENCES, TRIGGER'::text) pg_has_role(c.relowner, 'USAGE'::text)
or has_any_column_privilege(c.oid, 'SELECT, INSERT, UPDATE, REFERENCES'::text) or has_table_privilege(c.oid, 'SELECT, INSERT, UPDATE, DELETE, TRUNCATE, REFERENCES, TRIGGER'::text)
) or has_any_column_privilege(c.oid, 'SELECT, INSERT, UPDATE, REFERENCES'::text)
order by relname )
|] schema order by relname |]
return $ map tableFromRow rows
synonymousColumns :: [(Column,Column)] -> [Column] -> [[Column]] synonymousColumns :: [(Column,Column)] -> [Column] -> [[Column]]
synonymousColumns allSyns cols = synCols' synonymousColumns allSyns cols = synCols'
@@ -115,9 +172,9 @@ addForeignKeys rels = map addFk
addFk col = col { colFK = fk col } addFk col = col { colFK = fk col }
fk col = join $ relToFk col <$> find (lookupFn col) rels fk col = join $ relToFk col <$> find (lookupFn col) rels
lookupFn :: Column -> Relation -> Bool lookupFn :: Column -> Relation -> Bool
lookupFn c (Relation{relColumns=cs, relType=rty}) = c `elem` cs && rty==Child lookupFn c Relation{relColumns=cs, relType=rty} = c `elem` cs && rty==Child
-- lookupFn _ _ = False -- lookupFn _ _ = False
relToFk col (Relation{relColumns=cols, relFColumns=colsF}) = ForeignKey <$> colF relToFk col Relation{relColumns=cols, relFColumns=colsF} = ForeignKey <$> colF
where where
pos = elemIndex col cols pos = elemIndex col cols
colF = (colsF !!) <$> pos colF = (colsF !!) <$> pos
@@ -139,7 +196,7 @@ addManyToManyRelations rels = rels ++ addMirrorRelation (mapMaybe link2Relation
where where
links = join $ map (combinations 2) $ filter (not . null) $ groupWith groupFn $ filter ( (==Child). relType) rels links = join $ map (combinations 2) $ filter (not . null) $ groupWith groupFn $ filter ( (==Child). relType) rels
groupFn :: Relation -> Text groupFn :: Relation -> Text
groupFn (Relation{relTable=Table{tableSchema=s, tableName=t}}) = s<>"_"<>t groupFn Relation{relTable=Table{tableSchema=s, tableName=t}} = s<>"_"<>t
combinations k ns = filter ((k==).length) (subsequences ns) combinations k ns = filter ((k==).length) (subsequences ns)
addMirrorRelation [] = [] addMirrorRelation [] = []
addMirrorRelation (rel@(Relation t c ft fc _ lt lc1 lc2):rels') = Relation ft fc t c Many lt lc2 lc1 : rel : addMirrorRelation rels' addMirrorRelation (rel@(Relation t c ft fc _ lt lc1 lc2):rels') = Relation ft fc t c Many lt lc2 lc1 : rel : addMirrorRelation rels'
@@ -171,173 +228,170 @@ synonymousPrimaryKeys syns (key:keys) = key : newKeys ++ synonymousPrimaryKeys s
keySyns = filter ((\c -> colTable c == pkTable key && colName c == pkName key) . fst) syns keySyns = filter ((\c -> colTable c == pkTable key && colName c == pkName key) . fst) syns
newKeys = map ((\c -> PrimaryKey{pkTable=colTable c,pkName=colName c}) . snd) keySyns newKeys = map ((\c -> PrimaryKey{pkTable=colTable c,pkName=colName c}) . snd) keySyns
allTables :: H.Tx P.Postgres s [Table] allTables :: H.Query () [Table]
allTables = do allTables =
rows <- H.listEx $ [H.stmt| H.statement sql HE.unit decodeTables True
SELECT where
n.nspname AS table_schema, sql = [q|
c.relname AS table_name, SELECT
c.relkind = 'r' OR (c.relkind IN ('v','f')) n.nspname AS table_schema,
AND (pg_relation_is_updatable(c.oid::regclass, FALSE) & 8) = 8 c.relname AS table_name,
OR (EXISTS c.relkind = 'r' OR (c.relkind IN ('v','f'))
( SELECT 1 AND (pg_relation_is_updatable(c.oid::regclass, FALSE) & 8) = 8
FROM pg_trigger OR (EXISTS
WHERE pg_trigger.tgrelid = c.oid ( SELECT 1
AND (pg_trigger.tgtype::integer & 69) = 69) ) AS insertable FROM pg_trigger
FROM pg_class c WHERE pg_trigger.tgrelid = c.oid
JOIN pg_namespace n ON n.oid = c.relnamespace AND (pg_trigger.tgtype::integer & 69) = 69) ) AS insertable
WHERE c.relkind IN ('v','r','m') FROM pg_class c
AND n.nspname NOT IN ('pg_catalog', 'information_schema') JOIN pg_namespace n ON n.oid = c.relnamespace
GROUP BY table_schema, table_name, insertable WHERE c.relkind IN ('v','r','m')
ORDER BY table_schema, table_name AND n.nspname NOT IN ('pg_catalog', 'information_schema')
|] GROUP BY table_schema, table_name, insertable
return $ map tableFromRow rows ORDER BY table_schema, table_name |]
tableFromRow :: (Text, Text, Bool) -> Table allColumns :: [Table] -> H.Query () [Column]
tableFromRow (s, n, i) = Table s n i allColumns tabs =
H.statement sql HE.unit (decodeColumns tabs) True
allColumns :: [Table] -> H.Tx P.Postgres s [Column] where
allColumns tabs = do sql = [q|
cols <- H.listEx $ [H.stmt| SELECT DISTINCT
SELECT DISTINCT info.table_schema AS schema,
info.table_schema AS schema, info.table_name AS table_name,
info.table_name AS table_name, info.column_name AS name,
info.column_name AS name, info.ordinal_position AS position,
info.ordinal_position AS position, info.is_nullable::boolean AS nullable,
info.is_nullable::boolean AS nullable, info.data_type AS col_type,
info.data_type AS col_type, info.is_updatable::boolean AS updatable,
info.is_updatable::boolean AS updatable, info.character_maximum_length AS max_len,
info.character_maximum_length AS max_len, info.numeric_precision AS precision,
info.numeric_precision AS precision, info.column_default AS default_value,
info.column_default AS default_value, array_to_string(enum_info.vals, ',') AS enum
array_to_string(enum_info.vals, ',') AS enum FROM (
FROM ( /*
/* -- CTE based on information_schema.columns to remove the owner filter
-- CTE based on information_schema.columns to remove the owner filter */
*/ WITH columns AS (
WITH columns AS ( SELECT current_database()::information_schema.sql_identifier AS table_catalog,
SELECT current_database()::information_schema.sql_identifier AS table_catalog, nc.nspname::information_schema.sql_identifier AS table_schema,
nc.nspname::information_schema.sql_identifier AS table_schema, c.relname::information_schema.sql_identifier AS table_name,
c.relname::information_schema.sql_identifier AS table_name, a.attname::information_schema.sql_identifier AS column_name,
a.attname::information_schema.sql_identifier AS column_name, a.attnum::information_schema.cardinal_number AS ordinal_position,
a.attnum::information_schema.cardinal_number AS ordinal_position, pg_get_expr(ad.adbin, ad.adrelid)::information_schema.character_data AS column_default,
pg_get_expr(ad.adbin, ad.adrelid)::information_schema.character_data AS column_default, CASE
CASE WHEN a.attnotnull OR t.typtype = 'd'::"char" AND t.typnotnull THEN 'NO'::text
WHEN a.attnotnull OR t.typtype = 'd'::"char" AND t.typnotnull THEN 'NO'::text ELSE 'YES'::text
ELSE 'YES'::text END::information_schema.yes_or_no AS is_nullable,
END::information_schema.yes_or_no AS is_nullable, CASE
CASE WHEN t.typtype = 'd'::"char" THEN
WHEN t.typtype = 'd'::"char" THEN CASE
CASE WHEN bt.typelem <> 0::oid AND bt.typlen = (-1) THEN 'ARRAY'::text
WHEN bt.typelem <> 0::oid AND bt.typlen = (-1) THEN 'ARRAY'::text WHEN nbt.nspname = 'pg_catalog'::name THEN format_type(t.typbasetype, NULL::integer)
WHEN nbt.nspname = 'pg_catalog'::name THEN format_type(t.typbasetype, NULL::integer) ELSE 'USER-DEFINED'::text
ELSE 'USER-DEFINED'::text END
END ELSE
ELSE CASE
CASE WHEN t.typelem <> 0::oid AND t.typlen = (-1) THEN 'ARRAY'::text
WHEN t.typelem <> 0::oid AND t.typlen = (-1) THEN 'ARRAY'::text WHEN nt.nspname = 'pg_catalog'::name THEN format_type(a.atttypid, NULL::integer)
WHEN nt.nspname = 'pg_catalog'::name THEN format_type(a.atttypid, NULL::integer) ELSE 'USER-DEFINED'::text
ELSE 'USER-DEFINED'::text END
END END::information_schema.character_data AS data_type,
END::information_schema.character_data AS data_type, information_schema._pg_char_max_length(information_schema._pg_truetypid(a.*, t.*), information_schema._pg_truetypmod(a.*, t.*))::information_schema.cardinal_number AS character_maximum_length,
information_schema._pg_char_max_length(information_schema._pg_truetypid(a.*, t.*), information_schema._pg_truetypmod(a.*, t.*))::information_schema.cardinal_number AS character_maximum_length, information_schema._pg_char_octet_length(information_schema._pg_truetypid(a.*, t.*), information_schema._pg_truetypmod(a.*, t.*))::information_schema.cardinal_number AS character_octet_length,
information_schema._pg_char_octet_length(information_schema._pg_truetypid(a.*, t.*), information_schema._pg_truetypmod(a.*, t.*))::information_schema.cardinal_number AS character_octet_length, information_schema._pg_numeric_precision(information_schema._pg_truetypid(a.*, t.*), information_schema._pg_truetypmod(a.*, t.*))::information_schema.cardinal_number AS numeric_precision,
information_schema._pg_numeric_precision(information_schema._pg_truetypid(a.*, t.*), information_schema._pg_truetypmod(a.*, t.*))::information_schema.cardinal_number AS numeric_precision, information_schema._pg_numeric_precision_radix(information_schema._pg_truetypid(a.*, t.*), information_schema._pg_truetypmod(a.*, t.*))::information_schema.cardinal_number AS numeric_precision_radix,
information_schema._pg_numeric_precision_radix(information_schema._pg_truetypid(a.*, t.*), information_schema._pg_truetypmod(a.*, t.*))::information_schema.cardinal_number AS numeric_precision_radix, information_schema._pg_numeric_scale(information_schema._pg_truetypid(a.*, t.*), information_schema._pg_truetypmod(a.*, t.*))::information_schema.cardinal_number AS numeric_scale,
information_schema._pg_numeric_scale(information_schema._pg_truetypid(a.*, t.*), information_schema._pg_truetypmod(a.*, t.*))::information_schema.cardinal_number AS numeric_scale, information_schema._pg_datetime_precision(information_schema._pg_truetypid(a.*, t.*), information_schema._pg_truetypmod(a.*, t.*))::information_schema.cardinal_number AS datetime_precision,
information_schema._pg_datetime_precision(information_schema._pg_truetypid(a.*, t.*), information_schema._pg_truetypmod(a.*, t.*))::information_schema.cardinal_number AS datetime_precision, information_schema._pg_interval_type(information_schema._pg_truetypid(a.*, t.*), information_schema._pg_truetypmod(a.*, t.*))::information_schema.character_data AS interval_type,
information_schema._pg_interval_type(information_schema._pg_truetypid(a.*, t.*), information_schema._pg_truetypmod(a.*, t.*))::information_schema.character_data AS interval_type, NULL::integer::information_schema.cardinal_number AS interval_precision,
NULL::integer::information_schema.cardinal_number AS interval_precision, NULL::character varying::information_schema.sql_identifier AS character_set_catalog,
NULL::character varying::information_schema.sql_identifier AS character_set_catalog, NULL::character varying::information_schema.sql_identifier AS character_set_schema,
NULL::character varying::information_schema.sql_identifier AS character_set_schema, NULL::character varying::information_schema.sql_identifier AS character_set_name,
NULL::character varying::information_schema.sql_identifier AS character_set_name, CASE
CASE WHEN nco.nspname IS NOT NULL THEN current_database()
WHEN nco.nspname IS NOT NULL THEN current_database() ELSE NULL::name
ELSE NULL::name END::information_schema.sql_identifier AS collation_catalog,
END::information_schema.sql_identifier AS collation_catalog, nco.nspname::information_schema.sql_identifier AS collation_schema,
nco.nspname::information_schema.sql_identifier AS collation_schema, co.collname::information_schema.sql_identifier AS collation_name,
co.collname::information_schema.sql_identifier AS collation_name, CASE
CASE WHEN t.typtype = 'd'::"char" THEN current_database()
WHEN t.typtype = 'd'::"char" THEN current_database() ELSE NULL::name
ELSE NULL::name END::information_schema.sql_identifier AS domain_catalog,
END::information_schema.sql_identifier AS domain_catalog, CASE
CASE WHEN t.typtype = 'd'::"char" THEN nt.nspname
WHEN t.typtype = 'd'::"char" THEN nt.nspname ELSE NULL::name
ELSE NULL::name END::information_schema.sql_identifier AS domain_schema,
END::information_schema.sql_identifier AS domain_schema, CASE
CASE WHEN t.typtype = 'd'::"char" THEN t.typname
WHEN t.typtype = 'd'::"char" THEN t.typname ELSE NULL::name
ELSE NULL::name END::information_schema.sql_identifier AS domain_name,
END::information_schema.sql_identifier AS domain_name, current_database()::information_schema.sql_identifier AS udt_catalog,
current_database()::information_schema.sql_identifier AS udt_catalog, COALESCE(nbt.nspname, nt.nspname)::information_schema.sql_identifier AS udt_schema,
COALESCE(nbt.nspname, nt.nspname)::information_schema.sql_identifier AS udt_schema, COALESCE(bt.typname, t.typname)::information_schema.sql_identifier AS udt_name,
COALESCE(bt.typname, t.typname)::information_schema.sql_identifier AS udt_name, NULL::character varying::information_schema.sql_identifier AS scope_catalog,
NULL::character varying::information_schema.sql_identifier AS scope_catalog, NULL::character varying::information_schema.sql_identifier AS scope_schema,
NULL::character varying::information_schema.sql_identifier AS scope_schema, NULL::character varying::information_schema.sql_identifier AS scope_name,
NULL::character varying::information_schema.sql_identifier AS scope_name, NULL::integer::information_schema.cardinal_number AS maximum_cardinality,
NULL::integer::information_schema.cardinal_number AS maximum_cardinality, a.attnum::information_schema.sql_identifier AS dtd_identifier,
a.attnum::information_schema.sql_identifier AS dtd_identifier, 'NO'::character varying::information_schema.yes_or_no AS is_self_referencing,
'NO'::character varying::information_schema.yes_or_no AS is_self_referencing, 'NO'::character varying::information_schema.yes_or_no AS is_identity,
'NO'::character varying::information_schema.yes_or_no AS is_identity, NULL::character varying::information_schema.character_data AS identity_generation,
NULL::character varying::information_schema.character_data AS identity_generation, NULL::character varying::information_schema.character_data AS identity_start,
NULL::character varying::information_schema.character_data AS identity_start, NULL::character varying::information_schema.character_data AS identity_increment,
NULL::character varying::information_schema.character_data AS identity_increment, NULL::character varying::information_schema.character_data AS identity_maximum,
NULL::character varying::information_schema.character_data AS identity_maximum, NULL::character varying::information_schema.character_data AS identity_minimum,
NULL::character varying::information_schema.character_data AS identity_minimum, NULL::character varying::information_schema.yes_or_no AS identity_cycle,
NULL::character varying::information_schema.yes_or_no AS identity_cycle, 'NEVER'::character varying::information_schema.character_data AS is_generated,
'NEVER'::character varying::information_schema.character_data AS is_generated, NULL::character varying::information_schema.character_data AS generation_expression,
NULL::character varying::information_schema.character_data AS generation_expression, CASE
CASE WHEN c.relkind = 'r'::"char" OR (c.relkind = ANY (ARRAY['v'::"char", 'f'::"char"])) AND pg_column_is_updatable(c.oid::regclass, a.attnum, false) THEN 'YES'::text
WHEN c.relkind = 'r'::"char" OR (c.relkind = ANY (ARRAY['v'::"char", 'f'::"char"])) AND pg_column_is_updatable(c.oid::regclass, a.attnum, false) THEN 'YES'::text ELSE 'NO'::text
ELSE 'NO'::text END::information_schema.yes_or_no AS is_updatable
END::information_schema.yes_or_no AS is_updatable FROM pg_attribute a
FROM pg_attribute a LEFT JOIN pg_attrdef ad ON a.attrelid = ad.adrelid AND a.attnum = ad.adnum
LEFT JOIN pg_attrdef ad ON a.attrelid = ad.adrelid AND a.attnum = ad.adnum JOIN (pg_class c
JOIN (pg_class c JOIN pg_namespace nc ON c.relnamespace = nc.oid) ON a.attrelid = c.oid
JOIN pg_namespace nc ON c.relnamespace = nc.oid) ON a.attrelid = c.oid JOIN (pg_type t
JOIN (pg_type t JOIN pg_namespace nt ON t.typnamespace = nt.oid) ON a.atttypid = t.oid
JOIN pg_namespace nt ON t.typnamespace = nt.oid) ON a.atttypid = t.oid LEFT JOIN (pg_type bt
LEFT JOIN (pg_type bt JOIN pg_namespace nbt ON bt.typnamespace = nbt.oid) ON t.typtype = 'd'::"char" AND t.typbasetype = bt.oid
JOIN pg_namespace nbt ON bt.typnamespace = nbt.oid) ON t.typtype = 'd'::"char" AND t.typbasetype = bt.oid LEFT JOIN (pg_collation co
LEFT JOIN (pg_collation co JOIN pg_namespace nco ON co.collnamespace = nco.oid) ON a.attcollation = co.oid AND (nco.nspname <> 'pg_catalog'::name OR co.collname <> 'default'::name)
JOIN pg_namespace nco ON co.collnamespace = nco.oid) ON a.attcollation = co.oid AND (nco.nspname <> 'pg_catalog'::name OR co.collname <> 'default'::name) WHERE NOT pg_is_other_temp_schema(nc.oid) AND a.attnum > 0 AND NOT a.attisdropped AND (c.relkind = ANY (ARRAY['r'::"char", 'v'::"char", 'f'::"char"]))
WHERE NOT pg_is_other_temp_schema(nc.oid) AND a.attnum > 0 AND NOT a.attisdropped AND (c.relkind = ANY (ARRAY['r'::"char", 'v'::"char", 'f'::"char"])) /*--AND (pg_has_role(c.relowner, 'USAGE'::text) OR has_column_privilege(c.oid, a.attnum, 'SELECT, INSERT, UPDATE, REFERENCES'::text))*/
/*--AND (pg_has_role(c.relowner, 'USAGE'::text) OR has_column_privilege(c.oid, a.attnum, 'SELECT, INSERT, UPDATE, REFERENCES'::text))*/ )
) SELECT
SELECT table_schema,
table_schema, table_name,
table_name, column_name,
column_name, ordinal_position,
ordinal_position, is_nullable,
is_nullable, data_type,
data_type, is_updatable,
is_updatable, character_maximum_length,
character_maximum_length, numeric_precision,
numeric_precision, column_default,
column_default, udt_name
udt_name /*-- FROM information_schema.columns*/
/*-- FROM information_schema.columns*/ FROM columns
FROM columns WHERE table_schema NOT IN ('pg_catalog', 'information_schema')
WHERE table_schema NOT IN ('pg_catalog', 'information_schema') ) AS info
) AS info LEFT OUTER JOIN (
LEFT OUTER JOIN ( SELECT
SELECT n.nspname AS s,
n.nspname AS s, t.typname AS n,
t.typname AS n, array_agg(e.enumlabel ORDER BY e.enumsortorder) AS vals
array_agg(e.enumlabel ORDER BY e.enumsortorder) AS vals FROM pg_type t
FROM pg_type t JOIN pg_enum e ON t.oid = e.enumtypid
JOIN pg_enum e ON t.oid = e.enumtypid JOIN pg_catalog.pg_namespace n ON n.oid = t.typnamespace
JOIN pg_catalog.pg_namespace n ON n.oid = t.typnamespace GROUP BY s,n
GROUP BY s,n ) AS enum_info ON (info.udt_name = enum_info.n)
) AS enum_info ON (info.udt_name = enum_info.n) ORDER BY schema, position |]
ORDER BY schema, position
|]
return $ mapMaybe (columnFromRow tabs) cols
columnFromRow :: [Table] -> columnFromRow :: [Table] ->
(Text, Text, Text, (Text, Text, Text,
Int, Bool, Text, Int32, Bool, Text,
Bool, Maybe Int, Maybe Int, Bool, Maybe Int32, Maybe Int32,
Maybe Text, Maybe Text) Maybe Text, Maybe Text)
-> Maybe Column -> Maybe Column
columnFromRow tabs (s, t, n, pos, nul, typ, u, l, p, d, e) = buildColumn <$> table columnFromRow tabs (s, t, n, pos, nul, typ, u, l, p, d, e) = buildColumn <$> table
@@ -347,9 +401,11 @@ columnFromRow tabs (s, t, n, pos, nul, typ, u, l, p, d, e) = buildColumn <$> tab
parseEnum :: Maybe Text -> [Text] parseEnum :: Maybe Text -> [Text]
parseEnum str = fromMaybe [] $ split (==',') <$> str parseEnum str = fromMaybe [] $ split (==',') <$> str
allRelations :: [Table] -> [Column] -> H.Tx P.Postgres s [Relation] allRelations :: [Table] -> [Column] -> H.Query () [Relation]
allRelations tabs cols = do allRelations tabs cols =
rels <- H.listEx $ [H.stmt| H.statement sql HE.unit (decodeRelations tabs cols) True
where
sql = [q|
SELECT ns1.nspname AS table_schema, SELECT ns1.nspname AS table_schema,
tab.relname AS table_name, tab.relname AS table_name,
column_info.cols AS columns, column_info.cols AS columns,
@@ -373,9 +429,7 @@ allRelations tabs cols = do
LATERAL (SELECT * FROM pg_class WHERE pg_class.oid = confrelid) AS other, LATERAL (SELECT * FROM pg_class WHERE pg_class.oid = confrelid) AS other,
LATERAL (SELECT * FROM pg_namespace WHERE pg_namespace.oid = other.relnamespace) AS ns2 LATERAL (SELECT * FROM pg_namespace WHERE pg_namespace.oid = other.relnamespace) AS ns2
WHERE confrelid != 0 WHERE confrelid != 0
ORDER BY (conrelid, column_info.nums) ORDER BY (conrelid, column_info.nums) |]
|]
return $ mapMaybe (relationFromRow tabs cols) rels
relationFromRow :: [Table] -> [Column] -> (Text, Text, [Text], Text, Text, [Text]) -> Maybe Relation relationFromRow :: [Table] -> [Column] -> (Text, Text, [Text], Text, Text, [Text]) -> Maybe Relation
relationFromRow allTabs allCols (rs, rt, rcs, frs, frt, frcs) = relationFromRow allTabs allCols (rs, rt, rcs, frs, frt, frcs) =
@@ -388,119 +442,121 @@ relationFromRow allTabs allCols (rs, rt, rcs, frs, frt, frcs) =
cols = mapM (findCol rs rt) rcs cols = mapM (findCol rs rt) rcs
colsF = mapM (findCol frs frt) frcs colsF = mapM (findCol frs frt) frcs
allPrimaryKeys :: [Table] -> H.Tx P.Postgres s [PrimaryKey] allPrimaryKeys :: [Table] -> H.Query () [PrimaryKey]
allPrimaryKeys tabs = do allPrimaryKeys tabs =
pks <- H.listEx $ [H.stmt| H.statement sql HE.unit (decodePks tabs) True
/* where
-- CTE to replace information_schema.table_constraints to remove owner limit sql = [q|
*/ /*
WITH tc AS ( -- CTE to replace information_schema.table_constraints to remove owner limit
SELECT current_database()::information_schema.sql_identifier AS constraint_catalog, */
nc.nspname::information_schema.sql_identifier AS constraint_schema, WITH tc AS (
c.conname::information_schema.sql_identifier AS constraint_name, SELECT current_database()::information_schema.sql_identifier AS constraint_catalog,
current_database()::information_schema.sql_identifier AS table_catalog, nc.nspname::information_schema.sql_identifier AS constraint_schema,
nr.nspname::information_schema.sql_identifier AS table_schema, c.conname::information_schema.sql_identifier AS constraint_name,
r.relname::information_schema.sql_identifier AS table_name, current_database()::information_schema.sql_identifier AS table_catalog,
CASE c.contype nr.nspname::information_schema.sql_identifier AS table_schema,
WHEN 'c'::"char" THEN 'CHECK'::text r.relname::information_schema.sql_identifier AS table_name,
WHEN 'f'::"char" THEN 'FOREIGN KEY'::text CASE c.contype
WHEN 'p'::"char" THEN 'PRIMARY KEY'::text WHEN 'c'::"char" THEN 'CHECK'::text
WHEN 'u'::"char" THEN 'UNIQUE'::text WHEN 'f'::"char" THEN 'FOREIGN KEY'::text
ELSE NULL::text WHEN 'p'::"char" THEN 'PRIMARY KEY'::text
END::information_schema.character_data AS constraint_type, WHEN 'u'::"char" THEN 'UNIQUE'::text
CASE ELSE NULL::text
WHEN c.condeferrable THEN 'YES'::text END::information_schema.character_data AS constraint_type,
ELSE 'NO'::text CASE
END::information_schema.yes_or_no AS is_deferrable, WHEN c.condeferrable THEN 'YES'::text
CASE ELSE 'NO'::text
WHEN c.condeferred THEN 'YES'::text END::information_schema.yes_or_no AS is_deferrable,
ELSE 'NO'::text CASE
END::information_schema.yes_or_no AS initially_deferred WHEN c.condeferred THEN 'YES'::text
FROM pg_namespace nc, ELSE 'NO'::text
pg_namespace nr, END::information_schema.yes_or_no AS initially_deferred
pg_constraint c, FROM pg_namespace nc,
pg_class r pg_namespace nr,
WHERE nc.oid = c.connamespace AND nr.oid = r.relnamespace AND c.conrelid = r.oid AND (c.contype <> ALL (ARRAY['t'::"char", 'x'::"char"])) AND r.relkind = 'r'::"char" AND NOT pg_is_other_temp_schema(nr.oid) pg_constraint c,
/*--AND (pg_has_role(r.relowner, 'USAGE'::text) OR has_table_privilege(r.oid, 'INSERT, UPDATE, DELETE, TRUNCATE, REFERENCES, TRIGGER'::text) OR has_any_column_privilege(r.oid, 'INSERT, UPDATE, REFERENCES'::text))*/ pg_class r
UNION ALL WHERE nc.oid = c.connamespace AND nr.oid = r.relnamespace AND c.conrelid = r.oid AND (c.contype <> ALL (ARRAY['t'::"char", 'x'::"char"])) AND r.relkind = 'r'::"char" AND NOT pg_is_other_temp_schema(nr.oid)
SELECT current_database()::information_schema.sql_identifier AS constraint_catalog, /*--AND (pg_has_role(r.relowner, 'USAGE'::text) OR has_table_privilege(r.oid, 'INSERT, UPDATE, DELETE, TRUNCATE, REFERENCES, TRIGGER'::text) OR has_any_column_privilege(r.oid, 'INSERT, UPDATE, REFERENCES'::text))*/
nr.nspname::information_schema.sql_identifier AS constraint_schema, UNION ALL
(((((nr.oid::text || '_'::text) || r.oid::text) || '_'::text) || a.attnum::text) || '_not_null'::text)::information_schema.sql_identifier AS constraint_name, SELECT current_database()::information_schema.sql_identifier AS constraint_catalog,
current_database()::information_schema.sql_identifier AS table_catalog, nr.nspname::information_schema.sql_identifier AS constraint_schema,
nr.nspname::information_schema.sql_identifier AS table_schema, (((((nr.oid::text || '_'::text) || r.oid::text) || '_'::text) || a.attnum::text) || '_not_null'::text)::information_schema.sql_identifier AS constraint_name,
r.relname::information_schema.sql_identifier AS table_name, current_database()::information_schema.sql_identifier AS table_catalog,
'CHECK'::character varying::information_schema.character_data AS constraint_type, nr.nspname::information_schema.sql_identifier AS table_schema,
'NO'::character varying::information_schema.yes_or_no AS is_deferrable, r.relname::information_schema.sql_identifier AS table_name,
'NO'::character varying::information_schema.yes_or_no AS initially_deferred 'CHECK'::character varying::information_schema.character_data AS constraint_type,
FROM pg_namespace nr, 'NO'::character varying::information_schema.yes_or_no AS is_deferrable,
pg_class r, 'NO'::character varying::information_schema.yes_or_no AS initially_deferred
pg_attribute a FROM pg_namespace nr,
WHERE nr.oid = r.relnamespace AND r.oid = a.attrelid AND a.attnotnull AND a.attnum > 0 AND NOT a.attisdropped AND r.relkind = 'r'::"char" AND NOT pg_is_other_temp_schema(nr.oid) pg_class r,
/*--AND (pg_has_role(r.relowner, 'USAGE'::text) OR has_table_privilege(r.oid, 'INSERT, UPDATE, DELETE, TRUNCATE, REFERENCES, TRIGGER'::text) OR has_any_column_privilege(r.oid, 'INSERT, UPDATE, REFERENCES'::text))*/ pg_attribute a
), WHERE nr.oid = r.relnamespace AND r.oid = a.attrelid AND a.attnotnull AND a.attnum > 0 AND NOT a.attisdropped AND r.relkind = 'r'::"char" AND NOT pg_is_other_temp_schema(nr.oid)
/* /*--AND (pg_has_role(r.relowner, 'USAGE'::text) OR has_table_privilege(r.oid, 'INSERT, UPDATE, DELETE, TRUNCATE, REFERENCES, TRIGGER'::text) OR has_any_column_privilege(r.oid, 'INSERT, UPDATE, REFERENCES'::text))*/
-- CTE to replace information_schema.key_column_usage to remove owner limit ),
*/ /*
kc AS ( -- CTE to replace information_schema.key_column_usage to remove owner limit
SELECT current_database()::information_schema.sql_identifier AS constraint_catalog, */
ss.nc_nspname::information_schema.sql_identifier AS constraint_schema, kc AS (
ss.conname::information_schema.sql_identifier AS constraint_name, SELECT current_database()::information_schema.sql_identifier AS constraint_catalog,
current_database()::information_schema.sql_identifier AS table_catalog, ss.nc_nspname::information_schema.sql_identifier AS constraint_schema,
ss.nr_nspname::information_schema.sql_identifier AS table_schema, ss.conname::information_schema.sql_identifier AS constraint_name,
ss.relname::information_schema.sql_identifier AS table_name, current_database()::information_schema.sql_identifier AS table_catalog,
a.attname::information_schema.sql_identifier AS column_name, ss.nr_nspname::information_schema.sql_identifier AS table_schema,
(ss.x).n::information_schema.cardinal_number AS ordinal_position, ss.relname::information_schema.sql_identifier AS table_name,
CASE a.attname::information_schema.sql_identifier AS column_name,
WHEN ss.contype = 'f'::"char" THEN information_schema._pg_index_position(ss.conindid, ss.confkey[(ss.x).n]) (ss.x).n::information_schema.cardinal_number AS ordinal_position,
ELSE NULL::integer CASE
END::information_schema.cardinal_number AS position_in_unique_constraint WHEN ss.contype = 'f'::"char" THEN information_schema._pg_index_position(ss.conindid, ss.confkey[(ss.x).n])
FROM pg_attribute a, ELSE NULL::integer
( SELECT r.oid AS roid, END::information_schema.cardinal_number AS position_in_unique_constraint
r.relname, FROM pg_attribute a,
r.relowner, ( SELECT r.oid AS roid,
nc.nspname AS nc_nspname, r.relname,
nr.nspname AS nr_nspname, r.relowner,
c.oid AS coid, nc.nspname AS nc_nspname,
c.conname, nr.nspname AS nr_nspname,
c.contype, c.oid AS coid,
c.conindid, c.conname,
c.confkey, c.contype,
c.confrelid, c.conindid,
information_schema._pg_expandarray(c.conkey) AS x c.confkey,
FROM pg_namespace nr, c.confrelid,
pg_class r, information_schema._pg_expandarray(c.conkey) AS x
pg_namespace nc, FROM pg_namespace nr,
pg_constraint c pg_class r,
WHERE nr.oid = r.relnamespace AND r.oid = c.conrelid AND nc.oid = c.connamespace AND (c.contype = ANY (ARRAY['p'::"char", 'u'::"char", 'f'::"char"])) AND r.relkind = 'r'::"char" AND NOT pg_is_other_temp_schema(nr.oid)) ss pg_namespace nc,
WHERE ss.roid = a.attrelid AND a.attnum = (ss.x).x AND NOT a.attisdropped pg_constraint c
/*--AND (pg_has_role(ss.relowner, 'USAGE'::text) OR has_column_privilege(ss.roid, a.attnum, 'SELECT, INSERT, UPDATE, REFERENCES'::text))*/ WHERE nr.oid = r.relnamespace AND r.oid = c.conrelid AND nc.oid = c.connamespace AND (c.contype = ANY (ARRAY['p'::"char", 'u'::"char", 'f'::"char"])) AND r.relkind = 'r'::"char" AND NOT pg_is_other_temp_schema(nr.oid)) ss
) WHERE ss.roid = a.attrelid AND a.attnum = (ss.x).x AND NOT a.attisdropped
SELECT /*--AND (pg_has_role(ss.relowner, 'USAGE'::text) OR has_column_privilege(ss.roid, a.attnum, 'SELECT, INSERT, UPDATE, REFERENCES'::text))*/
kc.table_schema, )
kc.table_name, SELECT
kc.column_name kc.table_schema,
FROM kc.table_name,
/* kc.column_name
--information_schema.table_constraints tc, FROM
--information_schema.key_column_usage kc /*
*/ --information_schema.table_constraints tc,
tc, kc --information_schema.key_column_usage kc
WHERE */
tc.constraint_type = 'PRIMARY KEY' AND tc, kc
kc.table_name = tc.table_name AND WHERE
kc.table_schema = tc.table_schema AND tc.constraint_type = 'PRIMARY KEY' AND
kc.constraint_name = tc.constraint_name AND kc.table_name = tc.table_name AND
kc.table_schema NOT IN ('pg_catalog', 'information_schema') kc.table_schema = tc.table_schema AND
|] kc.constraint_name = tc.constraint_name AND
return $ mapMaybe (pkFromRow tabs) pks kc.table_schema NOT IN ('pg_catalog', 'information_schema') |]
pkFromRow :: [Table] -> (Schema, Text, Text) -> Maybe PrimaryKey pkFromRow :: [Table] -> (Schema, Text, Text) -> Maybe PrimaryKey
pkFromRow tabs (s, t, n) = PrimaryKey <$> table <*> pure n pkFromRow tabs (s, t, n) = PrimaryKey <$> table <*> pure n
where table = find (\tbl -> tableSchema tbl == s && tableName tbl == t) tabs where table = find (\tbl -> tableSchema tbl == s && tableName tbl == t) tabs
allSynonyms :: [Column] -> H.Tx P.Postgres s [(Column,Column)] allSynonyms :: [Column] -> H.Query () [(Column,Column)]
allSynonyms allCols = do allSynonyms cols =
syns <- H.listEx $ [H.stmt| H.statement sql HE.unit (decodeSynonyms cols) True
where
sql = [q|
WITH synonyms AS ( WITH synonyms AS (
/* /*
-- CTE to replace the view from information_schema because the information in it depended on the logged in role -- CTE to replace the view from information_schema because the information in it depended on the logged in role
@@ -562,9 +618,7 @@ allSynonyms allCols = do
syn_table_schema, syn_table_name, syn_table_schema, syn_table_name,
(regexp_matches(view_definition, CONCAT('\.', src_column_name, '\sAS\s("?)(.+?)\1(,|$)'), 'gn'))[2] AS syn_column_name /* " <- for syntax highlighting */ (regexp_matches(view_definition, CONCAT('\.', src_column_name, '\sAS\s("?)(.+?)\1(,|$)'), 'gn'))[2] AS syn_column_name /* " <- for syntax highlighting */
FROM synonyms FROM synonyms
) ) |]
|]
return $ mapMaybe (synonymFromRow allCols) syns
synonymFromRow :: [Column] -> (Text,Text,Text,Text,Text,Text) -> Maybe (Column,Column) synonymFromRow :: [Column] -> (Text,Text,Text,Text,Text,Text) -> Maybe (Column,Column)
synonymFromRow allCols (s1,t1,c1,s2,t2,c2) = (,) <$> col1 <*> col2 synonymFromRow allCols (s1,t1,c1,s2,t2,c2) = (,) <$> col1 <*> col2
+31 -27
View File
@@ -2,54 +2,58 @@
{-# LANGUAGE FlexibleInstances #-} {-# LANGUAGE FlexibleInstances #-}
{-# LANGUAGE TypeSynonymInstances #-} {-# LANGUAGE TypeSynonymInstances #-}
module PostgREST.Error (PgError, pgErrResponse, errResponse) where 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.Monoid ((<>))
import Data.String.Conversions (cs) import Data.String.Conversions (cs)
import Data.String.Utils (replace)
import Data.Text (Text) import Data.Text (Text)
import qualified Data.Text as T import qualified Data.Text as T
import qualified Hasql as H import qualified Hasql.Session as H
import qualified Hasql.Postgres 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)
type PgError = H.SessionError P.Postgres
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 :: PgError -> Response pgErrResponse :: H.Error -> 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 PgError where instance JSON.ToJSON H.Error where
toJSON (H.TxError (P.ErroneousResult 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),
"message" .= (cs m::T.Text), "message" .= (cs m::T.Text),
"details" .= (fmap cs d::Maybe T.Text), "details" .= (fmap cs d::Maybe T.Text),
"hint" .= (fmap cs h::Maybe T.Text)] "hint" .= (fmap cs h::Maybe T.Text)]
toJSON (H.TxError (P.NoResult d)) = JSON.object [ toJSON (H.ResultError (H.UnexpectedResult m)) = JSON.object [
"message" .= ("No response from server"::T.Text), "message" .= (cs m::T.Text)]
toJSON (H.ResultError (H.RowError i H.EndOfInput)) = JSON.object [
"message" .= ("Row error: end of input"::String),
"details" .=
("Attempt to parse more columns than there are in the result"::String),
"details" .= ("Row number " <> show i)]
toJSON (H.ResultError (H.RowError i H.UnexpectedNull)) = JSON.object [
"message" .= ("Row error: unexpected null"::String),
"details" .= ("Attempt to parse a NULL as some value."::String),
"details" .= ("Row number " <> show i)]
toJSON (H.ResultError (H.RowError i (H.ValueError d))) = JSON.object [
"message" .= ("Row error: Wrong value parser used"::String),
"details" .= d,
"details" .= ("Row number " <> show i)]
toJSON (H.ResultError (H.UnexpectedAmountOfRows i)) = JSON.object [
"message" .= ("Unexpected amount of rows"::String),
"details" .= i]
toJSON (H.ClientError d) = JSON.object [
"message" .= ("Database client error"::String),
"details" .= (fmap cs d::Maybe T.Text)] "details" .= (fmap cs d::Maybe T.Text)]
toJSON (H.TxError (P.UnexpectedResult m)) = JSON.object ["message" .= m]
toJSON (H.TxError P.NotInTransaction) = JSON.object [
"message" .= ("Not in transaction"::T.Text)]
toJSON (H.CxError (P.CantConnect d)) = JSON.object [
"message" .= ("Can't connect to the database"::T.Text),
"details" .= (fmap cs d::Maybe T.Text)]
toJSON (H.CxError (P.UnsupportedVersion v)) = JSON.object [
"message" .= ("Postgres version "++version++" is not supported") ]
where version = replace "0" "." (show v)
toJSON (H.ResultError m) = JSON.object ["message" .= m]
httpStatus :: PgError -> HT.Status httpStatus :: H.Error -> HT.Status
httpStatus (H.TxError (P.ErroneousResult codeBS _ _ _)) = httpStatus (H.ResultError (H.ServerError c _ _ _)) =
let code = cs codeBS in case cs c of
case code 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
@@ -75,5 +79,5 @@ httpStatus (H.TxError (P.ErroneousResult codeBS _ _ _)) =
"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.TxError (P.NoResult _)) = HT.status503 httpStatus (H.ResultError _) = HT.status500
httpStatus _ = HT.status500 httpStatus (H.ClientError _) = HT.status503
+41 -34
View File
@@ -9,21 +9,23 @@ import PostgREST.Config (AppConfig (..),
prettyVersion, prettyVersion,
readOptions) readOptions)
import PostgREST.DbStructure import PostgREST.DbStructure
import PostgREST.Error (PgError, pgErrResponse) import PostgREST.Error (errResponse, pgErrResponse)
import PostgREST.Middleware import PostgREST.Middleware
import PostgREST.QueryBuilder (inTransaction, Isolation(..))
import Control.Monad (unless, void) import Control.Monad (unless, void)
import Control.Monad.IO.Class (liftIO)
import Data.Aeson (encode)
import Data.Functor.Identity
import Data.Monoid ((<>)) import Data.Monoid ((<>))
import Data.Pool
import Data.String.Conversions (cs) import Data.String.Conversions (cs)
import Data.Text (Text)
import Data.Time.Clock.POSIX (getPOSIXTime) import Data.Time.Clock.POSIX (getPOSIXTime)
import qualified Hasql as H import qualified Hasql.Query as H
import qualified Hasql.Postgres as P import qualified Hasql.Connection as H
import qualified Hasql.Session as H
import qualified Hasql.Decoders as HD
import qualified Hasql.Encoders as HE
import qualified Network.HTTP.Types.Status as HT
import Network.Wai import Network.Wai
import Network.Wai.Handler.Warp hiding (Connection) import Network.Wai.Handler.Warp
import Network.Wai.Middleware.RequestLogger (logStdout) import Network.Wai.Middleware.RequestLogger (logStdout)
import System.IO (BufferMode (..), import System.IO (BufferMode (..),
hSetBuffering, stderr, hSetBuffering, stderr,
@@ -36,13 +38,14 @@ import Control.Concurrent (myThreadId)
import Control.Exception.Base (throwTo, AsyncException(..)) import Control.Exception.Base (throwTo, AsyncException(..))
#endif #endif
isServerVersionSupported :: H.Session P.Postgres IO Bool isServerVersionSupported :: H.Session Bool
isServerVersionSupported = do isServerVersionSupported = do
Identity (row :: Text) <- H.tx Nothing $ H.singleEx [H.stmt|SHOW server_version_num|] ver <- H.query () pgVersion
return $ read (cs row) >= minimumPgVersion return $ read (cs ver) >= minimumPgVersion
where
hasqlError :: PgError -> IO a pgVersion =
hasqlError = error . cs . encode H.statement "SHOW server_version_num"
HE.unit (HD.singleRow $ HD.value HD.text) True
main :: IO () main :: IO ()
main = do main = do
@@ -58,40 +61,44 @@ main = do
Prelude.putStrLn $ "Listening on port " ++ Prelude.putStrLn $ "Listening on port " ++
(show $ configPort conf :: String) (show $ configPort conf :: String)
let pgSettings = P.StringSettings $ cs (configDatabase conf) let pgSettings = cs (configDatabase conf)
appSettings = setPort port appSettings = setPort port
. setServerName (cs $ "postgrest/" <> prettyVersion) . setServerName (cs $ "postgrest/" <> prettyVersion)
$ defaultSettings $ defaultSettings
middle = logStdout . defaultMiddle middle = logStdout . defaultMiddle
poolSettings <- maybe (fail "Improper session settings") return $ pool <- createPool (H.acquire pgSettings)
H.poolSettings (fromIntegral $ configPool conf) 30 (either (const $ return ()) H.release) 1 1 (configPool conf)
pool :: H.Pool P.Postgres <- H.acquirePool pgSettings poolSettings
supportedOrError <- H.session pool isServerVersionSupported dbStructure <- withResource pool $ \case
either hasqlError Left err -> error $ show err
(\supported -> Right c -> do
unless supported $ supported <- H.run isServerVersionSupported c
error ( case supported of
"Cannot run in this PostgreSQL version, PostgREST needs at least " Left e -> error $ show e
<> show minimumPgVersion) Right good -> unless good $
) supportedOrError 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
H.releasePool pool destroyAllResources pool
throwTo tid UserInterrupt throwTo tid UserInterrupt
) Nothing ) Nothing
#endif #endif
let txSettings = Just (H.ReadCommitted, Just True)
dbOrError <- H.session pool $ H.tx txSettings $ getDbStructure (cs $ configSchema conf)
dbStructure <- either hasqlError return dbOrError
runSettings appSettings $ middle $ \ req respond -> do runSettings appSettings $ middle $ \ req respond -> do
time <- getPOSIXTime time <- getPOSIXTime
body <- strictRequestBody req body <- strictRequestBody req
resOrError <- liftIO $ H.session pool $ H.tx txSettings $ let handleReq = H.run $ inTransaction ReadCommitted
runWithClaims conf time (app dbStructure conf body) req (runWithClaims conf time (app dbStructure conf body) req)
either (respond . pgErrResponse) respond resOrError withResource pool $ \case
Left err -> respond $ errResponse HT.status500 (cs . show $ err)
Right c -> do
resOrError <- handleReq c
either (respond . pgErrResponse) respond resOrError
+6 -10
View File
@@ -7,8 +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 as H import qualified Hasql.Session as H
import qualified Hasql.Postgres as P
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)
@@ -25,28 +24,25 @@ import PostgREST.Error (errResponse)
import Prelude hiding(concat) import Prelude hiding(concat)
import qualified Data.Vector as V
import qualified Hasql.Backend as B
import qualified Data.Map.Lazy as M import qualified Data.Map.Lazy as M
runWithClaims :: forall s. AppConfig -> NominalDiffTime -> runWithClaims :: AppConfig -> NominalDiffTime ->
(Request -> H.Tx P.Postgres s Response) -> (Request -> H.Session Response) ->
Request -> H.Tx P.Postgres s Response Request -> H.Session Response
runWithClaims conf time app req = do runWithClaims conf time app req = do
_ <- H.unitEx $ stmt setAnon H.sql setAnon
case split (== ' ') (cs auth) of case split (== ' ') (cs auth) of
("Bearer" : tokenStr : _) -> ("Bearer" : tokenStr : _) ->
case jwtClaims jwtSecret tokenStr time of case jwtClaims jwtSecret tokenStr time of
Just claims -> Just claims ->
if M.member "role" claims if M.member "role" claims
then do then do
mapM_ H.unitEx $ stmt <$> claimsToSQL claims mapM_ H.sql $ claimsToSQL claims
app req app req
else invalidJWT else invalidJWT
_ -> invalidJWT _ -> invalidJWT
_ -> app req _ -> app req
where where
stmt c = B.Stmt c V.empty True
hdrs = requestHeaders req hdrs = requestHeaders req
jwtSecret = configJwtSecret conf jwtSecret = configJwtSecret conf
auth = fromMaybe "" $ lookup hAuthorization hdrs auth = fromMaybe "" $ lookup hAuthorization hdrs
+140 -87
View File
@@ -15,10 +15,10 @@ Any function that outputs a SQL fragment should be in this module.
module PostgREST.QueryBuilder ( module PostgREST.QueryBuilder (
addRelations addRelations
, addJoinConditions , addJoinConditions
, asJson
, callProc , callProc
, createReadStatement , createReadStatement
, createWriteStatement , createWriteStatement
, inTransaction
, operators , operators
, pgFmtIdent , pgFmtIdent
, pgFmtLit , pgFmtLit
@@ -26,28 +26,34 @@ module PostgREST.QueryBuilder (
, requestToCountQuery , requestToCountQuery
, sourceCTEName , sourceCTEName
, unquoted , unquoted
, ResultsWithCount
, Isolation(..)
) where ) where
import qualified Hasql as H import qualified Hasql.Query as H
import qualified Hasql.Backend as B import qualified Hasql.Session as H
import qualified Hasql.Postgres as P import qualified Hasql.Encoders as HE
import qualified Hasql.Decoders as HD
import qualified Data.Aeson as JSON import qualified Data.Aeson as JSON
import Data.Int (Int64)
import PostgREST.RangeQuery (NonnegRange, rangeLimit, rangeOffset) import PostgREST.RangeQuery (NonnegRange, rangeLimit, rangeOffset)
import Control.Error (note, fromMaybe, mapMaybe) import Control.Error (note, fromMaybe, mapMaybe)
import Data.Functor.Contravariant (contramap)
import qualified Data.HashMap.Strict as HM import qualified Data.HashMap.Strict as HM
import Data.List (find, (\\)) import Data.List (find, (\\))
import Data.Monoid ((<>)) import Data.Monoid ((<>))
import Data.Text (Text, intercalate, unwords, replace, isInfixOf, toLower, split) import Data.Text (Text, intercalate, unwords, replace, isInfixOf, toLower, split)
import qualified Data.Text as T (map, takeWhile) import qualified Data.Text as T (map, takeWhile)
import Data.String.Conversions (cs) import Data.String.Conversions (cs)
import Control.Applicative (empty, (<|>)) import Control.Applicative ((<|>))
import Control.Monad (join) import Control.Monad (join)
import Data.Tree (Tree(..)) import Data.Tree (Tree(..))
import qualified Data.Vector as V import qualified Data.Vector as V
import PostgREST.Types import PostgREST.Types
import qualified Data.Map as M import qualified Data.Map as M
import Text.InterpolatedString.Perl6 (qc)
import Text.Regex.TDFA ((=~)) import Text.Regex.TDFA ((=~))
import qualified Data.ByteString.Char8 as BS import qualified Data.ByteString.Char8 as BS
import Data.Scientific ( FPFormat (..) import Data.Scientific ( FPFormat (..)
@@ -57,70 +63,101 @@ import Data.Scientific ( FPFormat (..)
import Prelude hiding (unwords) import Prelude hiding (unwords)
import PostgREST.ApiRequest (PreferRepresentation (..)) import PostgREST.ApiRequest (PreferRepresentation (..))
type PStmt = H.Stmt P.Postgres
instance Monoid PStmt where
mappend (B.Stmt query params prep) (B.Stmt query' params' prep') =
B.Stmt (query <> query') (params <> params') (prep && prep')
mempty = B.Stmt "" empty True
type StatementT = PStmt -> PStmt
createReadStatement :: SqlQuery -> SqlQuery -> NonnegRange -> Bool -> Bool -> Bool -> B.Stmt P.Postgres {-| The generic query result format used by API responses -}
type ResultsWithCount = (Maybe Int64, Int64, BS.ByteString, BS.ByteString)
{-| Read and Write api requests use a similar response format which includes
various record counts and possible location header. This is the decoder
for that common type of query.
-}
decodeStandard :: HD.Result ResultsWithCount
decodeStandard =
HD.singleRow standardRow
where
standardRow = (,,,) <$> HD.nullableValue HD.int8 <*> HD.value HD.int8
<*> HD.value HD.bytea <*> HD.value HD.bytea
decodeStandardMay :: HD.Result (Maybe ResultsWithCount)
decodeStandardMay =
HD.maybeRow standardRow
where
standardRow = (,,,) <$> HD.nullableValue HD.int8 <*> HD.value HD.int8
<*> HD.value HD.bytea <*> HD.value HD.bytea
{-| JSON and CSV payloads from the client are given to us as
UniformObjects (objects who all have the same keys),
and we turn this into an old fasioned JSON array
-}
encodeUniformObjs :: HE.Params UniformObjects
encodeUniformObjs =
contramap (JSON.Array . V.map JSON.Object . unUniformObjects) (HE.value HE.json)
createReadStatement :: SqlQuery -> SqlQuery -> NonnegRange -> Bool -> Bool -> Bool ->
H.Query () ResultsWithCount
createReadStatement selectQuery countQuery range isSingle countTotal asCsv = createReadStatement selectQuery countQuery range isSingle countTotal asCsv =
B.Stmt ( H.statement sql HE.unit decodeStandard True
"WITH " <> sourceCTEName <> " AS (" <> selectQuery <> ") " <> where
"SELECT " <> intercalate ", " [ sql = [qc|
WITH {sourceCTEName} AS ({selectQuery}) SELECT {cols}
FROM ( SELECT * FROM {sourceCTEName} {limitF range}) t |]
countResultF = if countTotal then "("<>countQuery<>")" else "null"
cols = intercalate ", " [
countResultF <> " AS total_result_set", countResultF <> " AS total_result_set",
"pg_catalog.count(t) AS page_total", "pg_catalog.count(t) AS page_total",
"null AS header", "'' AS header",
bodyF <> " AS body" bodyF <> " AS body"
] <> ]
" FROM ( SELECT * FROM " <> sourceCTEName <> " " <> limitF range <> ") t" bodyF
) V.empty True | asCsv = asCsvF
where | isSingle = asJsonSingleF
countResultF = if countTotal then "("<>countQuery<>")" else "null" | otherwise = asJsonF
bodyF
| asCsv = asCsvF
| isSingle = asJsonSingleF
| otherwise = asJsonF
createWriteStatement :: QualifiedIdentifier -> SqlQuery -> SqlQuery -> Bool -> PreferRepresentation -> createWriteStatement :: QualifiedIdentifier -> SqlQuery -> SqlQuery -> Bool ->
[Text] -> Bool -> Payload -> B.Stmt P.Postgres PreferRepresentation -> [Text] -> Bool -> Payload ->
H.Query UniformObjects (Maybe ResultsWithCount)
createWriteStatement _ _ _ _ _ _ _ (PayloadParseError _) = undefined createWriteStatement _ _ _ _ _ _ _ (PayloadParseError _) = undefined
createWriteStatement _ _ mutateQuery _ None createWriteStatement _ _ mutateQuery _ None
_ _ (PayloadJSON (UniformObjects rows)) = _ _ (PayloadJSON (UniformObjects _)) =
B.Stmt ( H.statement sql encodeUniformObjs decodeStandardMay True
"WITH " <> sourceCTEName <> " AS (" <> mutateQuery <> ") " <> where
"SELECT null, 0, null, null" sql = [qc|
) (V.singleton . B.encodeValue . JSON.Array . V.map JSON.Object $ rows) True WITH {sourceCTEName} AS ({mutateQuery})
SELECT '', 0, '', '' |]
createWriteStatement qi _ mutateQuery isSingle HeadersOnly createWriteStatement qi _ mutateQuery isSingle HeadersOnly
pKeys _ (PayloadJSON (UniformObjects rows)) = pKeys _ (PayloadJSON (UniformObjects _)) =
B.Stmt ( H.statement sql encodeUniformObjs decodeStandardMay True
"WITH " <> sourceCTEName <> " AS (" <> mutateQuery <> " RETURNING " <> fromQi qi <> ".*" <> ") " <> where
"SELECT " <> intercalate ", " [ sql = [qc|
"null AS total_result_set", WITH {sourceCTEName} AS ({mutateQuery} RETURNING {fromQi qi}.*)
SELECT {cols}
FROM (SELECT 1 FROM {sourceCTEName}) t |]
cols = intercalate ", " [
"'' AS total_result_set",
"pg_catalog.count(t) AS page_total", "pg_catalog.count(t) AS page_total",
if isSingle then locationF pKeys else "null", if isSingle then locationF pKeys else "''",
"null" "''"
] <> ]
" FROM (SELECT 1 FROM " <> sourceCTEName <> ") t"
) (V.singleton . B.encodeValue . JSON.Array . V.map JSON.Object $ rows) True
createWriteStatement qi selectQuery mutateQuery isSingle Full createWriteStatement qi selectQuery mutateQuery isSingle Full
pKeys asCsv (PayloadJSON (UniformObjects rows)) = pKeys asCsv (PayloadJSON (UniformObjects _)) =
B.Stmt ( H.statement sql encodeUniformObjs decodeStandardMay True
"WITH " <> sourceCTEName <> " AS (" <> mutateQuery <> " RETURNING " <> fromQi qi <> ".*" <> ") " <> where
"SELECT " <> intercalate ", " [ sql = [qc|
"null AS total_result_set", -- when updateing it does not make sense WITH {sourceCTEName} AS ({mutateQuery} RETURNING {fromQi qi}.*)
SELECT {cols}
FROM ({selectQuery}) t |]
cols = intercalate ", " [
"'' AS total_result_set", -- when updateing it does not make sense
"pg_catalog.count(t) AS page_total", "pg_catalog.count(t) AS page_total",
if isSingle then locationF pKeys else "null" <> " AS header", if isSingle then locationF pKeys else "''" <> " AS header",
bodyF <> " AS body" bodyF <> " AS body"
] <> ]
" FROM ( "<>selectQuery<>") t" bodyF
) (V.singleton . B.encodeValue . JSON.Array . V.map JSON.Object $ rows) True | asCsv = asCsvF
where | isSingle = asJsonSingleF
bodyF | otherwise = asJsonF
| asCsv = asCsvF
| isSingle = asJsonSingleF
| otherwise = asJsonF
addRelations :: Schema -> [Relation] -> Maybe ReadRequest -> ReadRequest -> Either Text ReadRequest addRelations :: Schema -> [Relation] -> Maybe ReadRequest -> ReadRequest -> Either Text ReadRequest
addRelations schema allRelations parentNode node@(Node readNode@(query, (name, _)) forest) = addRelations schema allRelations parentNode node@(Node readNode@(query, (name, _)) forest) =
@@ -131,8 +168,8 @@ addRelations schema allRelations parentNode node@(Node readNode@(query, (name, _
$ findRelationByTable schema name parentTable $ findRelationByTable schema name parentTable
<|> findRelationByColumn schema parentTable name <|> findRelationByColumn schema parentTable name
addRel :: (ReadQuery, (NodeName, Maybe Relation)) -> Relation -> (ReadQuery, (NodeName, Maybe Relation)) addRel :: (ReadQuery, (NodeName, Maybe Relation)) -> Relation -> (ReadQuery, (NodeName, Maybe Relation))
addRel (q, (n, _)) r = (q {from=fromRelation}, (n, Just r)) addRel (query', (n, _)) r = (query' {from=fromRelation}, (n, Just r))
where fromRelation = map (\t -> if t == n then tableName (relTable r) else t) (from q) where fromRelation = map (\t -> if t == n then tableName (relTable r) else t) (from query')
_ -> Node (query, (name, Nothing)) <$> updatedForest _ -> Node (query, (name, Nothing)) <$> updatedForest
where where
@@ -150,13 +187,13 @@ addJoinConditions :: Schema -> ReadRequest -> Either Text ReadRequest
addJoinConditions schema (Node (query, (n, r)) forest) = addJoinConditions schema (Node (query, (n, r)) forest) =
case r of case r of
Nothing -> Node (updatedQuery, (n,r)) <$> updatedForest -- this is the root node Nothing -> Node (updatedQuery, (n,r)) <$> updatedForest -- this is the root node
Just rel@(Relation{relType=Child}) -> Node (addCond updatedQuery (getJoinConditions rel),(n,r)) <$> updatedForest Just rel@Relation{relType=Child} -> Node (addCond updatedQuery (getJoinConditions rel),(n,r)) <$> updatedForest
Just (Relation{relType=Parent}) -> Node (updatedQuery, (n,r)) <$> updatedForest Just Relation{relType=Parent} -> Node (updatedQuery, (n,r)) <$> updatedForest
Just rel@(Relation{relType=Many, relLTable=(Just linkTable)}) -> Just rel@Relation{relType=Many, relLTable=(Just linkTable)} ->
Node (qq, (n, r)) <$> updatedForest Node (qq, (n, r)) <$> updatedForest
where where
q = addCond updatedQuery (getJoinConditions rel) query' = addCond updatedQuery (getJoinConditions rel)
qq = q{from=tableName linkTable : from q} qq = query'{from=tableName linkTable : from query'}
_ -> Left "unknown relation" _ -> Left "unknown relation"
where where
-- add parentTable and parentJoinConditions to the query -- add parentTable and parentJoinConditions to the query
@@ -164,23 +201,23 @@ addJoinConditions schema (Node (query, (n, r)) forest) =
where where
parentJoinConditions = map (getJoinConditions . snd) parents parentJoinConditions = map (getJoinConditions . snd) parents
parents = mapMaybe (getParents . rootLabel) forest parents = mapMaybe (getParents . rootLabel) forest
getParents (_, (tbl, Just rel@(Relation{relType=Parent}))) = Just (tbl, rel) getParents (_, (tbl, Just rel@Relation{relType=Parent})) = Just (tbl, rel)
getParents _ = Nothing getParents _ = Nothing
updatedForest = mapM (addJoinConditions schema) forest updatedForest = mapM (addJoinConditions schema) forest
addCond q con = q{flt_=con ++ flt_ q} addCond query' con = query'{flt_=con ++ flt_ query'}
asJson :: StatementT callProc :: QualifiedIdentifier -> JSON.Object -> H.Query () (Maybe JSON.Value)
asJson s = s { callProc qi params =
B.stmtTemplate = H.statement sql HE.unit decodeObj True
"array_to_json(coalesce(array_agg(row_to_json(t)), '{}'))::character varying from ("
<> B.stmtTemplate s <> ") t" }
callProc :: QualifiedIdentifier -> JSON.Object -> PStmt
callProc qi params = do
let args = intercalate "," $ map assignment (HM.toList params)
B.Stmt ("select * from " <> fromQi qi <> "(" <> args <> ")") empty True
where where
assignment (n,v) = pgFmtIdent n <> ":=" <> insertableValue v sql = [qc| SELECT array_to_json(
coalesce(array_agg(row_to_json(t)), '\{}')
)::character varying
from ({_callSql}) t |]
_args = intercalate "," $ map _assignment (HM.toList params)
_assignment (n,v) = pgFmtIdent n <> ":=" <> insertableValue v
_callSql = [qc| select * from {fromQi qi}({_args}) |] :: BS.ByteString
decodeObj = HD.maybeRow (HD.value HD.json)
operators :: [(Text, SqlFragment)] operators :: [(Text, SqlFragment)]
operators = [ operators = [
@@ -222,8 +259,8 @@ requestToCountQuery schema (DbRead (Node (Select _ _ conditions _, (mainTbl, _))
("WHERE " <> intercalate " AND " ( map (pgFmtCondition (QualifiedIdentifier schema mainTbl)) localConditions )) `emptyOnNull` localConditions ("WHERE " <> intercalate " AND " ( map (pgFmtCondition (QualifiedIdentifier schema mainTbl)) localConditions )) `emptyOnNull` localConditions
] ]
where where
fn (Filter{value=VText _}) = True fn Filter{value=VText _} = True
fn (Filter{value=VForeignKey _ _}) = False fn Filter{value=VForeignKey _ _} = False
localConditions = filter fn conditions localConditions = filter fn conditions
requestToQuery :: Schema -> DbRequest -> SqlQuery requestToQuery :: Schema -> DbRequest -> SqlQuery
@@ -268,19 +305,19 @@ requestToQuery schema (DbRead (Node (Select colSelects tbls conditions ord, (nod
filterParentConditions parentTable (Filter _ _ (VForeignKey (QualifiedIdentifier "" t) _)) = parentTable == t filterParentConditions parentTable (Filter _ _ (VForeignKey (QualifiedIdentifier "" t) _)) = parentTable == t
filterParentConditions _ _ = False filterParentConditions _ _ = False
getQueryParts :: Tree ReadNode -> ([(SqlFragment, TableName)], [SqlFragment]) -> ([(SqlFragment,TableName)], [SqlFragment]) getQueryParts :: Tree ReadNode -> ([(SqlFragment, TableName)], [SqlFragment]) -> ([(SqlFragment,TableName)], [SqlFragment])
getQueryParts (Node n@(_, (name, Just (Relation {relType=Child,relTable=Table{tableName=table}}))) forst) (j,s) = (j,sel:s) getQueryParts (Node n@(_, (name, Just Relation{relType=Child,relTable=Table{tableName=table}})) forst) (j,s) = (j,sel:s)
where where
sel = "COALESCE((" sel = "COALESCE(("
<> "SELECT array_to_json(array_agg(row_to_json("<>pgFmtIdent table<>"))) " <> "SELECT array_to_json(array_agg(row_to_json("<>pgFmtIdent table<>"))) "
<> "FROM (" <> subquery <> ") " <> pgFmtIdent table <> "FROM (" <> subquery <> ") " <> pgFmtIdent table
<> "), '[]') AS " <> pgFmtIdent name <> "), '[]') AS " <> pgFmtIdent name
where subquery = requestToQuery schema (DbRead (Node n forst)) where subquery = requestToQuery schema (DbRead (Node n forst))
getQueryParts (Node n@(_, (name, Just (Relation {relType=Parent,relTable=Table{tableName=table}}))) forst) (j,s) = (joi:j,sel:s) getQueryParts (Node n@(_, (name, Just Relation{relType=Parent,relTable=Table{tableName=table}})) forst) (j,s) = (joi:j,sel:s)
where where
sel = "row_to_json(" <> pgFmtIdent table <> ".*) AS "<>pgFmtIdent name --TODO must be singular sel = "row_to_json(" <> pgFmtIdent table <> ".*) AS "<>pgFmtIdent name --TODO must be singular
joi = ("( " <> subquery <> " ) AS " <> pgFmtIdent table, table) joi = ("( " <> subquery <> " ) AS " <> pgFmtIdent table, table)
where subquery = requestToQuery schema (DbRead (Node n forst)) where subquery = requestToQuery schema (DbRead (Node n forst))
getQueryParts (Node n@(_, (name, Just (Relation {relType=Many,relTable=Table{tableName=table}}))) forst) (j,s) = (j,sel:s) getQueryParts (Node n@(_, (name, Just Relation{relType=Many,relTable=Table{tableName=table}})) forst) (j,s) = (j,sel:s)
where where
sel = "COALESCE ((" sel = "COALESCE (("
<> "SELECT array_to_json(array_agg(row_to_json("<>pgFmtIdent table<>"))) " <> "SELECT array_to_json(array_agg(row_to_json("<>pgFmtIdent table<>"))) "
@@ -299,7 +336,7 @@ requestToQuery schema (DbMutate (Insert mainTbl (PayloadJSON (UniformObjects row
"INSERT INTO ", fromQi qi, "INSERT INTO ", fromQi qi,
" (" <> colsString <> ")" <> " (" <> colsString <> ")" <>
" SELECT " <> colsString <> " SELECT " <> colsString <>
" FROM json_populate_recordset(null::" , fromQi qi, ", ?)" " FROM json_populate_recordset(null::" , fromQi qi, ", $1)"
] ]
requestToQuery schema (DbMutate (Update mainTbl (PayloadJSON (UniformObjects rows)) conditions)) = requestToQuery schema (DbMutate (Update mainTbl (PayloadJSON (UniformObjects rows)) conditions)) =
case rows V.!? 0 of case rows V.!? 0 of
@@ -339,7 +376,7 @@ asCsvF :: SqlFragment
asCsvF = asCsvHeaderF <> " || '\n' || " <> asCsvBodyF asCsvF = asCsvHeaderF <> " || '\n' || " <> asCsvBodyF
where where
asCsvHeaderF = asCsvHeaderF =
"(SELECT string_agg(a.k, ',')" <> "(SELECT coalesce(string_agg(a.k, ','), '')" <>
" FROM (" <> " FROM (" <>
" SELECT json_object_keys(r)::TEXT as k" <> " SELECT json_object_keys(r)::TEXT as k" <>
" FROM ( " <> " FROM ( " <>
@@ -350,10 +387,10 @@ asCsvF = asCsvHeaderF <> " || '\n' || " <> asCsvBodyF
asCsvBodyF = "coalesce(string_agg(substring(t::text, 2, length(t::text) - 2), '\n'), '')" asCsvBodyF = "coalesce(string_agg(substring(t::text, 2, length(t::text) - 2), '\n'), '')"
asJsonF :: SqlFragment asJsonF :: SqlFragment
asJsonF = "array_to_json(array_agg(row_to_json(t)))::character varying" asJsonF = "coalesce(array_to_json(array_agg(row_to_json(t))), '[]')::character varying"
asJsonSingleF :: SqlFragment --TODO! unsafe when the query actually returns multiple rows, used only on inserting and returning single element asJsonSingleF :: SqlFragment --TODO! unsafe when the query actually returns multiple rows, used only on inserting and returning single element
asJsonSingleF = "string_agg(row_to_json(t)::text, ',')::character varying " asJsonSingleF = "coalesce(string_agg(row_to_json(t)::text, ','), '')::character varying "
locationF :: [Text] -> SqlFragment locationF :: [Text] -> SqlFragment
locationF pKeys = locationF pKeys =
@@ -365,8 +402,7 @@ locationF pKeys =
if null pKeys if null pKeys
then "" then ""
else " WHERE json_data.key IN ('" <> intercalate "','" pKeys <> "')" else " WHERE json_data.key IN ('" <> intercalate "','" pKeys <> "')"
) <> ) <> ")"
")"
limitF :: NonnegRange -> SqlFragment limitF :: NonnegRange -> SqlFragment
limitF r = "LIMIT " <> limit <> " OFFSET " <> offset limitF r = "LIMIT " <> limit <> " OFFSET " <> offset
@@ -468,3 +504,20 @@ 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"
+6 -6
View File
@@ -24,7 +24,7 @@ import Data.Maybe (fromMaybe, listToMaybe)
import Prelude import Prelude
type NonnegRange = Range Int type NonnegRange = Range Integer
rangeParse :: BS.ByteString -> NonnegRange rangeParse :: BS.ByteString -> NonnegRange
rangeParse range = do rangeParse range = do
@@ -41,28 +41,28 @@ rangeParse range = do
rangeRequested :: RequestHeaders -> NonnegRange rangeRequested :: RequestHeaders -> NonnegRange
rangeRequested = rangeParse . fromMaybe "" . lookup hRange rangeRequested = rangeParse . fromMaybe "" . lookup hRange
restrictRange :: Maybe Int -> NonnegRange -> NonnegRange restrictRange :: Maybe Integer -> NonnegRange -> NonnegRange
restrictRange Nothing r = r restrictRange Nothing r = r
restrictRange (Just limit) r = restrictRange (Just limit) r =
rangeIntersection r $ rangeIntersection r $
Range BoundaryBelowAll (BoundaryAbove $ rangeOffset r + limit - 1) Range BoundaryBelowAll (BoundaryAbove $ rangeOffset r + limit - 1)
rangeLimit :: NonnegRange -> Maybe Int rangeLimit :: NonnegRange -> Maybe Integer
rangeLimit range = rangeLimit range =
case [rangeLower range, rangeUpper range] of case [rangeLower range, rangeUpper range] of
[BoundaryBelow from, BoundaryAbove to] -> Just (1 + to - from) [BoundaryBelow from, BoundaryAbove to] -> Just (1 + to - from)
_ -> Nothing _ -> Nothing
rangeOffset :: NonnegRange -> Int rangeOffset :: NonnegRange -> Integer
rangeOffset range = rangeOffset range =
case rangeLower range of case rangeLower range of
BoundaryBelow from -> from BoundaryBelow from -> from
_ -> error "range without lower bound" -- should never happen _ -> error "range without lower bound" -- should never happen
rangeGeq :: Int -> NonnegRange rangeGeq :: Integer -> NonnegRange
rangeGeq n = rangeGeq n =
Range (BoundaryBelow n) BoundaryAboveAll Range (BoundaryBelow n) BoundaryAboveAll
rangeLeq :: Int -> NonnegRange rangeLeq :: Integer -> NonnegRange
rangeLeq n = rangeLeq n =
Range BoundaryBelowAll (BoundaryAbove n) Range BoundaryBelowAll (BoundaryAbove n)
+7 -3
View File
@@ -5,6 +5,7 @@ import qualified Data.ByteString.Lazy as BL
import qualified Data.ByteString as BS import qualified Data.ByteString as BS
import qualified Data.Vector as V import qualified Data.Vector as V
import Data.Aeson import Data.Aeson
import Data.Int (Int32)
data DbStructure = DbStructure { data DbStructure = DbStructure {
dbTables :: [Table] dbTables :: [Table]
@@ -31,12 +32,12 @@ data Column =
Column { Column {
colTable :: Table colTable :: Table
, colName :: Text , colName :: Text
, colPosition :: Int , colPosition :: Int32
, colNullable :: Bool , colNullable :: Bool
, colType :: Text , colType :: Text
, colUpdatable :: Bool , colUpdatable :: Bool
, colMaxLen :: Maybe Int , colMaxLen :: Maybe Int32
, colPrecision :: Maybe Int , colPrecision :: Maybe Int32
, colDefault :: Maybe Text , colDefault :: Maybe Text
, colEnum :: [Text] , colEnum :: [Text]
, colFK :: Maybe ForeignKey , colFK :: Maybe ForeignKey
@@ -90,6 +91,9 @@ data Relation = Relation {
newtype UniformObjects = UniformObjects (V.Vector Object) newtype UniformObjects = UniformObjects (V.Vector Object)
deriving (Show, Eq) deriving (Show, Eq)
unUniformObjects :: UniformObjects -> V.Vector Object
unUniformObjects (UniformObjects objs) = objs
-- | When Hasql supports the COPY command then we can -- | When Hasql supports the COPY command then we can
-- have a special payload just for CSV, but until -- have a special payload just for CSV, but until
-- then CSV is converted to a JSON array. -- then CSV is converted to a JSON array.
+3 -2
View File
@@ -2,6 +2,7 @@ flags: {}
packages: packages:
- '.' - '.'
extra-deps: extra-deps:
- hasql-0.19.3.3
- Ranged-sets-0.3.0 - Ranged-sets-0.3.0
- packdeps-0.4.1 - packdeps-0.4.2.1
resolver: nightly-2015-10-27 resolver: lts-5.0
+3 -5
View File
@@ -5,16 +5,14 @@ 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 Hasql as H
import Hasql.Postgres as P
import SpecHelper import SpecHelper
import PostgREST.Types (DbStructure(..)) import PostgREST.Types (DbStructure(..))
-- }}} -- }}}
spec :: DbStructure -> H.Pool P.Postgres -> Spec spec :: DbStructure -> H.Connection -> Spec
spec struct pool = around (withApp cfgDefault struct pool) spec struct c = around (withApp cfgDefault struct c)
$ describe "authorization" $ do $ describe "authorization" $ do
it "hides tables that anonymous does not own" $ it "hides tables that anonymous does not own" $
+3 -5
View File
@@ -5,9 +5,7 @@ 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 Hasql as H
import Hasql.Postgres as P
import SpecHelper import SpecHelper
import PostgREST.Types (DbStructure(..)) import PostgREST.Types (DbStructure(..))
@@ -15,8 +13,8 @@ import PostgREST.Types (DbStructure(..))
import Network.HTTP.Types import Network.HTTP.Types
-- }}} -- }}}
spec :: DbStructure -> H.Pool P.Postgres -> Spec spec :: DbStructure -> H.Connection -> Spec
spec struct pool = around (withApp cfgDefault struct pool) $ describe "CORS" $ do spec struct c = around (withApp cfgDefault struct c) $ describe "CORS" $ do
let preflightHeaders = [ let preflightHeaders = [
("Accept", "*/*"), ("Accept", "*/*"),
("Origin", "http://example.com"), ("Origin", "http://example.com"),
+4 -6
View File
@@ -4,17 +4,15 @@ import Test.Hspec
import Test.Hspec.Wai import Test.Hspec.Wai
import Text.Heredoc import Text.Heredoc
import Hasql as H
import Hasql.Postgres as P
import SpecHelper import SpecHelper
import PostgREST.Types (DbStructure(..)) import PostgREST.Types (DbStructure(..))
import qualified Hasql.Connection as H
import Network.HTTP.Types import Network.HTTP.Types
spec :: DbStructure -> H.Pool P.Postgres -> Spec spec :: DbStructure -> H.Connection -> Spec
spec struct pool = beforeAll resetDb spec struct c = beforeAll resetDb
. around (withApp cfgDefault struct pool) $ . 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 -5
View File
@@ -5,9 +5,6 @@ import Test.Hspec.Wai
import Test.Hspec.Wai.JSON import Test.Hspec.Wai.JSON
import Network.Wai.Test (SResponse(simpleBody,simpleHeaders,simpleStatus)) import Network.Wai.Test (SResponse(simpleBody,simpleHeaders,simpleStatus))
import Hasql as H
import Hasql.Postgres as P
import SpecHelper import SpecHelper
import PostgREST.Types (DbStructure(..)) import PostgREST.Types (DbStructure(..))
@@ -17,11 +14,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(..))
spec :: DbStructure -> H.Pool P.Postgres -> Spec spec :: DbStructure -> H.Connection -> Spec
spec struct pool = beforeAll_ resetDb $ around (withApp cfgDefault struct pool) $ do spec struct c = beforeAll_ resetDb $ around (withApp cfgDefault struct c) $ do
describe "Posting new record" $ do describe "Posting new record" $ do
context "disparate csv types" $ do context "disparate csv types" $ do
it "accepts disparate json types" $ do it "accepts disparate json types" $ do
+4 -6
View File
@@ -5,17 +5,15 @@ 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 Hasql as H
import Hasql.Postgres as P
import SpecHelper import SpecHelper
import PostgREST.Types (DbStructure(..)) import PostgREST.Types (DbStructure(..))
spec :: DbStructure -> H.Pool P.Postgres -> Spec spec :: DbStructure -> H.Connection -> Spec
spec struct pool = spec struct c =
beforeAll resetDb beforeAll resetDb
. around (withApp (cfgLimitRows 3) struct pool) $ . 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 -5
View File
@@ -5,16 +5,14 @@ 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 Hasql as H
import Hasql.Postgres as P
import SpecHelper import SpecHelper
import PostgREST.Types (DbStructure(..)) import PostgREST.Types (DbStructure(..))
import Text.Heredoc import Text.Heredoc
spec :: DbStructure -> H.Pool P.Postgres -> Spec spec :: DbStructure -> H.Connection -> Spec
spec struct pool = around (withApp cfgDefault struct pool) $ do spec struct c = around (withApp cfgDefault struct c) $ 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 -6
View File
@@ -5,16 +5,14 @@ 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 Hasql as H
import Hasql.Postgres as P
import SpecHelper import SpecHelper
import PostgREST.Types (DbStructure(..)) import PostgREST.Types (DbStructure(..))
spec :: DbStructure -> H.Pool P.Postgres -> Spec spec :: DbStructure -> H.Connection -> Spec
spec struct pool = beforeAll resetDb spec struct c = beforeAll resetDb
. around (withApp cfgDefault struct pool) $ . around (withApp cfgDefault struct c) $
describe "GET /items" $ do describe "GET /items" $ do
context "without range headers" $ do context "without range headers" $ do
+3 -5
View File
@@ -3,17 +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 Hasql as H
import Hasql.Postgres as P
import SpecHelper import SpecHelper
import PostgREST.Types (DbStructure(..)) import PostgREST.Types (DbStructure(..))
import Network.HTTP.Types import Network.HTTP.Types
spec :: DbStructure -> H.Pool P.Postgres -> Spec spec :: DbStructure -> H.Connection -> Spec
spec struct pool = around (withApp cfgDefault struct pool) $ do 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 "/" [] ""
+22 -16
View File
@@ -3,7 +3,11 @@ module Main where
import Test.Hspec import Test.Hspec
import SpecHelper import SpecHelper
--import PostgREST.Types (DbStructure(..)) import qualified Hasql.Session as H
import qualified Hasql.Connection as H
import PostgREST.DbStructure (getDbStructure)
import Data.String.Conversions (cs)
import qualified Feature.AuthSpec import qualified Feature.AuthSpec
import qualified Feature.CorsSpec import qualified Feature.CorsSpec
@@ -18,20 +22,22 @@ main :: IO ()
main = do main = do
setupDb setupDb
pool <- specDbPool H.acquire (cs dbString) >>= \case
dbStructure <- specDbStructure pool Left err -> error $ show err
Right c -> do
-- Not using hspec-discover because we want to precompute dbOrErr <- H.run (getDbStructure "test") c
-- the db structure and pass it to specs for speed -- Not using hspec-discover because we want to precompute
hspec $ specs dbStructure pool -- the db structure and pass it to specs for speed
either (error.show) (hspec . specs c) dbOrErr
H.release c
where where
specs dbStructure pool = do specs conn dbStructure = do
describe "Feature.AuthSpec" $ Feature.AuthSpec.spec dbStructure pool describe "Feature.AuthSpec" $ Feature.AuthSpec.spec dbStructure conn
describe "Feature.CorsSpec" $ Feature.CorsSpec.spec dbStructure pool describe "Feature.CorsSpec" $ Feature.CorsSpec.spec dbStructure conn
describe "Feature.DeleteSpec" $ Feature.DeleteSpec.spec dbStructure pool describe "Feature.DeleteSpec" $ Feature.DeleteSpec.spec dbStructure conn
describe "Feature.InsertSpec" $ Feature.InsertSpec.spec dbStructure pool describe "Feature.InsertSpec" $ Feature.InsertSpec.spec dbStructure conn
describe "Feature.QueryLimitedSpec" $ Feature.QueryLimitedSpec.spec dbStructure pool describe "Feature.QueryLimitedSpec" $ Feature.QueryLimitedSpec.spec dbStructure conn
describe "Feature.QuerySpec" $ Feature.QuerySpec.spec dbStructure pool describe "Feature.QuerySpec" $ Feature.QuerySpec.spec dbStructure conn
describe "Feature.RangeSpec" $ Feature.RangeSpec.spec dbStructure pool describe "Feature.RangeSpec" $ Feature.RangeSpec.spec dbStructure conn
describe "Feature.StructureSpec" $ Feature.StructureSpec.spec dbStructure pool describe "Feature.StructureSpec" $ Feature.StructureSpec.spec dbStructure conn
+16 -46
View File
@@ -2,16 +2,8 @@ module SpecHelper where
import Network.Wai import Network.Wai
import Test.Hspec import Test.Hspec
import Test.Hspec.Wai
import Hasql as H
import Hasql.Backend as B
import Hasql.Postgres as P
import Data.String.Conversions (cs) import Data.String.Conversions (cs)
import Data.Monoid
import Data.Text hiding (map)
import qualified Data.Vector as V
import Data.Time.Clock.POSIX (getPOSIXTime) import Data.Time.Clock.POSIX (getPOSIXTime)
import Control.Monad (void) import Control.Monad (void)
@@ -19,57 +11,47 @@ import Network.HTTP.Types.Header (Header, ByteRange, renderByteRange,
hRange, hAuthorization, hAccept) hRange, hAuthorization, hAccept)
import Codec.Binary.Base64.String (encode) import Codec.Binary.Base64.String (encode)
import Data.CaseInsensitive (CI(..)) import Data.CaseInsensitive (CI(..))
import Data.Maybe (fromMaybe)
import Text.Regex.TDFA ((=~)) import Text.Regex.TDFA ((=~))
import qualified Data.ByteString.Char8 as BS 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.App (app)
import PostgREST.Config (AppConfig(..)) import PostgREST.Config (AppConfig(..))
import PostgREST.Middleware import PostgREST.Middleware
import PostgREST.Error(pgErrResponse) import PostgREST.Error(pgErrResponse)
import PostgREST.DbStructure
import PostgREST.Types import PostgREST.Types
import PostgREST.QueryBuilder (inTransaction, Isolation(..))
dbString :: String dbString :: String
dbString = "postgres://postgrest_test_authenticator@localhost:5432/postgrest_test" dbString = "postgres://postgrest_test_authenticator@localhost:5432/postgrest_test"
cfg :: String -> Maybe Int -> AppConfig cfg :: String -> Maybe Integer -> AppConfig
cfg conStr = AppConfig conStr 3000 "postgrest_test_anonymous" "test" (secret "safe") 10 cfg conStr = AppConfig conStr 3000 "postgrest_test_anonymous" "test" (secret "safe") 10
cfgDefault :: AppConfig cfgDefault :: AppConfig
cfgDefault = cfg dbString Nothing cfgDefault = cfg dbString Nothing
cfgLimitRows :: Int -> AppConfig cfgLimitRows :: Integer -> AppConfig
cfgLimitRows = cfg dbString . Just cfgLimitRows = cfg dbString . Just
testPoolOpts :: PoolSettings withApp :: AppConfig -> DbStructure -> H.Connection
testPoolOpts = fromMaybe (error "bad settings") $ H.poolSettings 1 30
pgSettings :: P.Settings
pgSettings = P.StringSettings $ cs dbString
specDbPool :: IO (H.Pool P.Postgres)
specDbPool = H.acquirePool pgSettings testPoolOpts
specDbStructure :: H.Pool P.Postgres -> IO DbStructure
specDbStructure pool = do
dbOrError <- H.session pool $ H.tx specTxSettings
$ getDbStructure "test"
either (fail . show) return dbOrError
withApp :: AppConfig -> DbStructure -> H.Pool P.Postgres
-> ActionWith Application -> IO () -> ActionWith Application -> IO ()
withApp config dbStructure pool perform = do withApp config dbStructure c perform = do
perform $ middle $ \req resp -> do perform $ defaultMiddle $ \req resp -> do
time <- getPOSIXTime time <- getPOSIXTime
body <- strictRequestBody req body <- strictRequestBody req
result <- liftIO $ H.session pool $ H.tx specTxSettings let handleReq = H.run $ inTransaction ReadCommitted
$ runWithClaims config time (app dbStructure config body) req (runWithClaims config time (app dbStructure config body) req)
either (resp . pgErrResponse) resp result
where middle = defaultMiddle 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
@@ -106,15 +88,3 @@ authHeaderBasic u p =
authHeaderJWT :: String -> Header authHeaderJWT :: String -> Header
authHeaderJWT token = authHeaderJWT token =
(hAuthorization, cs $ "Bearer " ++ token) (hAuthorization, cs $ "Bearer " ++ token)
testPool :: IO(H.Pool P.Postgres)
testPool = H.acquirePool pgSettings testPoolOpts
clearTable :: Text -> IO ()
clearTable table = do
pool <- testPool
void . liftIO $ H.session pool $ H.tx Nothing $
H.unitEx $ B.Stmt ("truncate table test." <> table <> " cascade") V.empty True
specTxSettings :: Maybe (TxIsolationLevel, Maybe Bool)
specTxSettings = Just (H.ReadCommitted, Just True)