Merge pull request #620 from ruslantalpa/rpc_refactor

Rpc refactor
This commit is contained in:
Joe Nelson
2016-07-06 20:51:05 -07:00
committed by GitHub
7 changed files with 120 additions and 69 deletions
+1 -1
View File
@@ -8,7 +8,7 @@ This project adheres to [Semantic Versioning](http://semver.org/).
### Added ### Added
- Ability to generate an OpenAPI spec - @mainx07, @hudayou, @ruslantalpa, @begriffs - Ability to generate an OpenAPI spec - @mainx07, @hudayou, @ruslantalpa, @begriffs
- Ability to set addresses to listen on - @hudayou - Ability to set addresses to listen on - @hudayou
- Filtering, shaping and embedding with &select for the /rpc path - @ruslantalpa
- Output names of used-defined types (instead of 'USER-DEFINED') - @martingms - Output names of used-defined types (instead of 'USER-DEFINED') - @martingms
### Fixed ### Fixed
+28 -18
View File
@@ -14,7 +14,7 @@ import Data.List (find, delete)
import Data.Maybe (fromMaybe, fromJust, mapMaybe) import Data.Maybe (fromMaybe, fromJust, mapMaybe)
import Data.Ranged.Ranges (emptyRange) import Data.Ranged.Ranges (emptyRange)
import Data.String.Conversions (cs) import Data.String.Conversions (cs)
import Data.Text (Text, replace, strip) import Data.Text (Text, replace, strip, isInfixOf, dropWhile, drop)
import Data.Tree import Data.Tree
import qualified Hasql.Pool as P import qualified Hasql.Pool as P
@@ -62,7 +62,7 @@ import PostgREST.QueryBuilder ( callProc
import PostgREST.Types import PostgREST.Types
import PostgREST.OpenAPI import PostgREST.OpenAPI
import Prelude import Prelude hiding (dropWhile, drop)
postgrest :: AppConfig -> IORef DbStructure -> P.Pool -> Application postgrest :: AppConfig -> IORef DbStructure -> P.Pool -> Application
@@ -186,22 +186,22 @@ app dbStructure conf apiRequest =
(ActionInvoke, TargetProc qi, (ActionInvoke, TargetProc qi,
Just (PayloadJSON (UniformObjects payload))) -> do Just (PayloadJSON (UniformObjects payload))) -> do
exists <- H.query qi doesProcExist let p = V.head payload
if exists singular = iPreferSingular apiRequest
then do jwtSecret = configJwtSecret conf
let p = V.head payload returnType = lookup (qiName qi) $ dbProcs dbStructure
jwtSecret = configJwtSecret conf returnsJWT = fromMaybe False $ isInfixOf "jwt_claims" <$> returnType
respondToRange $ do case readSqlParts of
row <- H.query () (callProc qi p topLevelRange shouldCount) Left e -> return $ responseLBS status400 [jsonH] $ cs e
returnJWT <- H.query qi doesProcReturnJWT Right (q,cq) -> respondToRange $ do
row <- H.query () (callProc qi p q cq topLevelRange shouldCount singular)
let (tableTotal, queryTotal, body) = fromMaybe (Just 0, 0, emptyArray) row let (tableTotal, queryTotal, body) = fromMaybe (Just 0, 0, emptyArray) row
(status, contentRange) = rangeHeader queryTotal tableTotal (status, contentRange) = rangeHeader queryTotal tableTotal
in in
return $ responseLBS status [jsonH, contentRange] return $ responseLBS status [jsonH, contentRange]
(if returnJWT (if returnsJWT
then "{\"token\":\"" <> cs (tokenJWT jwtSecret body) <> "\"}" then "{\"token\":\"" <> cs (tokenJWT jwtSecret body) <> "\"}"
else cs $ encode body) else cs $ encode body)
else return notFound
(ActionRead, TargetRoot, Nothing) -> do (ActionRead, TargetRoot, Nothing) -> do
let encodeApi ti = encodeOpenAPI ti host port let encodeApi ti = encodeOpenAPI ti host port
@@ -241,7 +241,7 @@ app dbStructure conf apiRequest =
schema = cs $ configSchema conf schema = cs $ configSchema conf
shouldCount = iPreferCount apiRequest shouldCount = iPreferCount apiRequest
topLevelRange = fromMaybe allRange $ M.lookup "limit" $ iRange apiRequest topLevelRange = fromMaybe allRange $ M.lookup "limit" $ iRange apiRequest
readDbRequest = DbRead <$> buildReadRequest (configMaxRows conf) (dbRelations dbStructure) apiRequest readDbRequest = DbRead <$> buildReadRequest (configMaxRows conf) (dbRelations dbStructure) (dbProcs dbStructure) apiRequest
mutateDbRequest = DbMutate <$> buildMutateRequest apiRequest mutateDbRequest = DbMutate <$> buildMutateRequest apiRequest
selectQuery = requestToQuery schema False <$> readDbRequest selectQuery = requestToQuery schema False <$> readDbRequest
countQuery = requestToCountQuery schema <$> readDbRequest countQuery = requestToCountQuery schema <$> readDbRequest
@@ -326,9 +326,10 @@ addFiltersOrdersRanges apiRequest = foldr1 (liftA2 (.)) [
filters = mapM pRequestFilter flts filters = mapM pRequestFilter flts
where where
action = iAction apiRequest action = iAction apiRequest
flts = if action == ActionRead flts
then iFilters apiRequest | action == ActionRead = iFilters apiRequest
else filter (( '.' `elem` ) . fst) $ iFilters apiRequest -- there can be no filters on the root table whre we are doing insert/update | action == ActionInvoke = iFilters apiRequest
| otherwise = filter (( '.' `elem` ) . fst) $ iFilters apiRequest -- there can be no filters on the root table whre we are doing insert/update
orders :: Either ParseError [(Path, [OrderTerm])] orders :: Either ParseError [(Path, [OrderTerm])]
orders = mapM pRequestOrder $ iOrder apiRequest orders = mapM pRequestOrder $ iOrder apiRequest
ranges :: Either ParseError [(Path, NonnegRange)] ranges :: Either ParseError [(Path, NonnegRange)]
@@ -340,8 +341,8 @@ treeRestrictRange maxRows_ request = pure $ nodeRestrictRange maxRows_ `fmap` re
nodeRestrictRange :: Maybe Integer -> ReadNode -> ReadNode nodeRestrictRange :: Maybe Integer -> ReadNode -> ReadNode
nodeRestrictRange m (q@Select {range_=r}, i) = (q{range_=restrictRange m r }, i) nodeRestrictRange m (q@Select {range_=r}, i) = (q{range_=restrictRange m r }, i)
buildReadRequest :: Maybe Integer -> [Relation] -> ApiRequest -> Either Text ReadRequest buildReadRequest :: Maybe Integer -> [Relation] -> [(Text, Text)] -> ApiRequest -> Either Text ReadRequest
buildReadRequest maxRows allRels apiRequest = buildReadRequest maxRows allRels allProcs apiRequest =
treeRestrictRange maxRows =<< treeRestrictRange maxRows =<<
augumentRequestWithJoin schema relations =<< augumentRequestWithJoin schema relations =<<
first formatParserError readRequest first formatParserError readRequest
@@ -350,6 +351,14 @@ buildReadRequest maxRows allRels apiRequest =
let target = iTarget apiRequest in let target = iTarget apiRequest in
case target of case target of
(TargetIdent (QualifiedIdentifier s t) ) -> Just (s, t) (TargetIdent (QualifiedIdentifier s t) ) -> Just (s, t)
(TargetProc (QualifiedIdentifier s p) ) -> Just (s, t)
where
returnType = fromMaybe "" $ lookup p allProcs
-- we are looking for results looking like "SETOF schema.tablename" and want to extract tablename
t = if "SETOF " `isInfixOf` returnType
then drop 1 $ dropWhile (/= '.') returnType
else p
_ -> Nothing _ -> Nothing
action :: Action action :: Action
@@ -369,6 +378,7 @@ buildReadRequest maxRows allRels apiRequest =
ActionCreate -> fakeSourceRelations ++ allRels ActionCreate -> fakeSourceRelations ++ allRels
ActionUpdate -> fakeSourceRelations ++ allRels ActionUpdate -> fakeSourceRelations ++ allRels
ActionDelete -> fakeSourceRelations ++ allRels ActionDelete -> fakeSourceRelations ++ allRels
ActionInvoke -> fakeSourceRelations ++ allRels
_ -> allRels _ -> allRels
where fakeSourceRelations = mapMaybe (toSourceRelation rootTableName) allRels -- see comment in toSourceRelation where fakeSourceRelations = mapMaybe (toSourceRelation rootTableName) allRels -- see comment in toSourceRelation
+11 -33
View File
@@ -6,8 +6,6 @@
module PostgREST.DbStructure ( module PostgREST.DbStructure (
getDbStructure getDbStructure
, accessibleTables , accessibleTables
, doesProcExist
, doesProcReturnJWT
) where ) where
import qualified Hasql.Decoders as HD import qualified Hasql.Decoders as HD
@@ -16,7 +14,6 @@ import qualified Hasql.Query as H
import Control.Applicative import Control.Applicative
import Control.Monad (join, replicateM) import Control.Monad (join, replicateM)
import Data.Functor.Contravariant (contramap)
import Data.List (elemIndex, find, sort, import Data.List (elemIndex, find, sort,
subsequences, transpose) subsequences, transpose)
import Data.Maybe (fromJust, fromMaybe, isJust, import Data.Maybe (fromJust, fromMaybe, isJust,
@@ -38,6 +35,7 @@ getDbStructure schema = do
syns <- H.query () $ allSynonyms cols syns <- H.query () $ allSynonyms cols
rels <- H.query () $ allRelations tabs cols rels <- H.query () $ allRelations tabs cols
keys <- H.query () $ allPrimaryKeys tabs keys <- H.query () $ allPrimaryKeys tabs
procs <- H.query schema accessibleProcs
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
@@ -48,13 +46,9 @@ getDbStructure schema = do
, dbColumns = cols' , dbColumns = cols'
, dbRelations = rels' , dbRelations = rels'
, dbPrimaryKeys = keys' , dbPrimaryKeys = keys'
, dbProcs = procs
} }
encodeQi :: HE.Params QualifiedIdentifier
encodeQi =
contramap qiSchema (HE.value HE.text) <>
contramap qiName (HE.value HE.text)
decodeTables :: HD.Result [Table] decodeTables :: HD.Result [Table]
decodeTables = decodeTables =
HD.rowsList tblRow HD.rowsList tblRow
@@ -104,32 +98,16 @@ decodeSynonyms cols =
<*> HD.value HD.text <*> HD.value HD.text <*> 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 accessibleProcs :: H.Query Schema [(Text, Text)]
doesProcExist = accessibleProcs =
H.statement sql encodeQi (HD.singleRow (HD.value HD.bool)) True H.statement sql (HE.value HE.text) (HD.rowsList ((,) <$> HD.value HD.text <*> HD.value HD.text)) True
where where
sql = [q| SELECT EXISTS ( sql = [q|
SELECT 1 SELECT p.proname as "proc_name", pg_get_function_result(p.oid) as "return_type"
FROM pg_catalog.pg_namespace n FROM pg_namespace n
JOIN pg_catalog.pg_proc p JOIN pg_proc p
ON pronamespace = n.oid ON pronamespace = n.oid
WHERE nspname = $1 WHERE n.nspname = $1|]
AND proname = $2
) |]
doesProcReturnJWT :: H.Query QualifiedIdentifier Bool
doesProcReturnJWT =
H.statement sql encodeQi (HD.singleRow (HD.value HD.bool)) True
where
sql = [q| SELECT EXISTS (
SELECT 1
FROM pg_catalog.pg_namespace n
JOIN pg_catalog.pg_proc p
ON pronamespace = n.oid
WHERE nspname = $1
AND proname = $2
AND pg_catalog.pg_get_function_result(p.oid) like '%jwt_claims'
) |]
accessibleTables :: H.Query Schema [Table] accessibleTables :: H.Query Schema [Table]
accessibleTables = accessibleTables =
+26 -13
View File
@@ -203,29 +203,39 @@ addJoinConditions schema (Node nn@(query, (n, r, a)) forest) =
addCond query' con = query'{flt_=con ++ flt_ query'} addCond query' con = query'{flt_=con ++ flt_ query'}
type ProcResults = (Maybe Int64, Int64, JSON.Value) type ProcResults = (Maybe Int64, Int64, JSON.Value)
callProc :: QualifiedIdentifier -> JSON.Object -> NonnegRange -> Bool -> H.Query () (Maybe ProcResults) callProc :: QualifiedIdentifier -> JSON.Object -> SqlQuery -> SqlQuery -> NonnegRange -> Bool -> Bool -> H.Query () (Maybe ProcResults)
callProc qi params range countTotal = callProc qi params selectQuery countQuery _ countTotal isSingle =
unicodeStatement sql HE.unit decodeProc True unicodeStatement sql HE.unit decodeProc True
where where
sql = [qc| sql = [qc|
WITH t AS (select * {_callSql}) WITH {sourceCTEName} AS ({_callSql})
SELECT SELECT
{_countExpr} as countTotal, {countResultF} AS total_result_set,
pg_catalog.count(1) as countResult, pg_catalog.count(t) AS page_total,
array_to_json( case
coalesce(array_agg(row_to_json(r)), '\{}') when pg_catalog.count(1) > 1 then
)::character varying {bodyF}
FROM (select * from t {limitF range}) r; else
coalesce(((array_agg(row_to_json(t)))[1]->{_procName})::character varying, {bodyF})
end as body
FROM ({selectQuery}) t;
|] |]
-- FROM (select * from {sourceCTEName} {limitF range}) t;
countResultF = if countTotal then "("<>countQuery<>")" else "null::bigint" :: Text
_args = intercalate "," $ map _assignment (HM.toList params) _args = intercalate "," $ map _assignment (HM.toList params)
_procName = pgFmtLit $ qiName qi
_assignment (n,v) = pgFmtIdent n <> ":=" <> insertableValue v _assignment (n,v) = pgFmtIdent n <> ":=" <> insertableValue v
_callSql = [qc| from {fromQi qi}({_args}) |] :: Text _callSql = [qc|select * from {fromQi qi}({_args}) |] :: Text
_countExpr = if countTotal _countExpr = if countTotal
then "(select pg_catalog.count(1) from t)" then [qc|(select pg_catalog.count(1) from {sourceCTEName})|]
else "null::bigint" :: Text else "null::bigint" :: Text
decodeProc = HD.maybeRow procRow decodeProc = HD.maybeRow procRow
procRow = (,,) <$> HD.nullableValue HD.int8 <*> HD.value HD.int8 procRow = (,,) <$> HD.nullableValue HD.int8 <*> HD.value HD.int8
<*> HD.value HD.json <*> HD.value HD.json
bodyF
| isSingle = asJsonSingleF
| otherwise = asJsonF
operators :: [(Text, SqlFragment)] operators :: [(Text, SqlFragment)]
operators = [ operators = [
@@ -263,10 +273,13 @@ requestToCountQuery _ (DbMutate _) = undefined
requestToCountQuery schema (DbRead (Node (Select _ _ conditions _ _, (mainTbl, _, _)) _)) = requestToCountQuery schema (DbRead (Node (Select _ _ conditions _ _, (mainTbl, _, _)) _)) =
unwords [ unwords [
"SELECT pg_catalog.count(1)", "SELECT pg_catalog.count(1)",
"FROM ", fromQi $ QualifiedIdentifier schema mainTbl, "FROM ", fromQi qi,
("WHERE " <> intercalate " AND " ( map (pgFmtCondition (QualifiedIdentifier schema mainTbl)) localConditions )) `emptyOnNull` localConditions ("WHERE " <> intercalate " AND " ( map (pgFmtCondition qi) localConditions )) `emptyOnNull` localConditions
] ]
where where
qi = if mainTbl == sourceCTEName
then QualifiedIdentifier "" mainTbl
else QualifiedIdentifier schema mainTbl
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
+1
View File
@@ -13,6 +13,7 @@ data DbStructure = DbStructure {
, dbColumns :: [Column] , dbColumns :: [Column]
, dbRelations :: [Relation] , dbRelations :: [Relation]
, dbPrimaryKeys :: [PrimaryKey] , dbPrimaryKeys :: [PrimaryKey]
, dbProcs :: [(Text,Text)]
} deriving (Show, Eq) } deriving (Show, Eq)
type Schema = Text type Schema = Text
+41 -4
View File
@@ -454,6 +454,38 @@ spec = do
post "/rpc/getitemrange" [json| { "min": 2, "max": 4 } |] `shouldRespondWith` post "/rpc/getitemrange" [json| { "min": 2, "max": 4 } |] `shouldRespondWith`
[json| [ {"id": 3}, {"id":4} ] |] [json| [ {"id": 3}, {"id":4} ] |]
context "shaping the response returned by a proc" $ do
it "returns a project" $
post "/rpc/getproject" [json| { "id": 1} |] `shouldRespondWith`
[json|[{"id":1,"name":"Windows 7","client_id":1}]|]
it "can filter proc results" $
post "/rpc/getallprojects?id=gt.1&id=lt.5&select=id" [json| {} |] `shouldRespondWith`
[json|[{"id":2},{"id":3},{"id":4}]|]
it "can limit proc results" $
post "/rpc/getallprojects?id=gt.1&id=lt.5&select=id?limit=2&offset=1" [json| {} |]
`shouldRespondWith` ResponseMatcher {
matchBody = Just [json|[{"id":3},{"id":4}]|]
, matchStatus = 206
, matchHeaders = ["Content-Range" <:> "1-2/3"]
}
it "prefer singular" $
request methodPost "/rpc/getproject"
[("Prefer","plurality=singular")] [json| { "id": 1} |] `shouldRespondWith`
[json|{"id":1,"name":"Windows 7","client_id":1}|]
it "select works on the first level" $
post "/rpc/getproject?select=id,name" [json| { "id": 1} |] `shouldRespondWith`
[json|[{"id":1,"name":"Windows 7"}]|]
it "can embed foreign entities to the items returned by a proc" $
post "/rpc/getproject?select=id,name,client{id},tasks{id}" [json| { "id": 1} |] `shouldRespondWith`
[json|[{"id":1,"name":"Windows 7","client":{"id":1},"tasks":[{"id":1},{"id":2}]}]|]
context "a proc that returns an empty rowset" $ context "a proc that returns an empty rowset" $
it "returns empty json array" $ it "returns empty json array" $
post "/rpc/test_empty_rowset" [json| {} |] `shouldRespondWith` post "/rpc/test_empty_rowset" [json| {} |] `shouldRespondWith`
@@ -462,11 +494,11 @@ spec = do
context "a proc that returns plain text" $ do context "a proc that returns plain text" $ do
it "returns proper json" $ it "returns proper json" $
post "/rpc/sayhello" [json| { "name": "world" } |] `shouldRespondWith` post "/rpc/sayhello" [json| { "name": "world" } |] `shouldRespondWith`
[json| [{"sayhello":"Hello, world"}] |] [json|"Hello, world"|]
it "can handle unicode" $ it "can handle unicode" $
post "/rpc/sayhello" [json| { "name": "" } |] `shouldRespondWith` post "/rpc/sayhello" [json| { "name": "" } |] `shouldRespondWith`
[json| [{"sayhello":"Hello, ¥"}] |] [json|"Hello, ¥"|]
context "improper input" $ do context "improper input" $ do
it "rejects unknown content type even if payload is good" $ it "rejects unknown content type even if payload is good" $
@@ -477,6 +509,11 @@ spec = do
request methodPost "/rpc/sayhello" request methodPost "/rpc/sayhello"
(acceptHdrs "application/json") "sdfsdf" (acceptHdrs "application/json") "sdfsdf"
`shouldRespondWith` 400 `shouldRespondWith` 400
-- it used to be 404 and it makes sense but in another part we decided that it's good to return
-- PostgreSQL errors (and have the proxy handle them) and this saves us an aditional query on each rpc request
it "responds with 400 on an unexisting proc" $
post "/rpc/fake" [json| {} |] `shouldRespondWith` 400
context "unsupported verbs" $ do context "unsupported verbs" $ do
it "DELETE fails" $ it "DELETE fails" $
@@ -497,9 +534,9 @@ spec = do
it "executes the proc exactly once per request" $ do it "executes the proc exactly once per request" $ do
post "/rpc/callcounter" [json| {} |] `shouldRespondWith` post "/rpc/callcounter" [json| {} |] `shouldRespondWith`
[json| [{"callcounter":1}] |] [json|1|]
post "/rpc/callcounter" [json| {} |] `shouldRespondWith` post "/rpc/callcounter" [json| {} |] `shouldRespondWith`
[json| [{"callcounter":2}] |] [json|2|]
describe "weird requests" $ do describe "weird requests" $ do
it "can query as normal" $ do it "can query as normal" $ do
+12
View File
@@ -1021,6 +1021,18 @@ create table orders (
shipping_address_id int references addresses(id) shipping_address_id int references addresses(id)
); );
CREATE FUNCTION getproject(id int) RETURNS SETOF projects
LANGUAGE sql
AS $_$
SELECT * FROM test.projects WHERE id = $1;
$_$;
CREATE FUNCTION getallprojects() RETURNS SETOF projects
LANGUAGE sql
AS $_$
SELECT * FROM test.projects;
$_$;
-- --
-- PostgreSQL database dump complete -- PostgreSQL database dump complete
-- --