From 2f166088c803c8bbf974ffecae1ceb2af2d7ed1a Mon Sep 17 00:00:00 2001 From: Ruslan Talpa Date: Fri, 27 May 2016 11:18:52 +0300 Subject: [PATCH] response shaping and filtering for rpc proc calls --- src/PostgREST/App.hs | 32 ++++++++++++++++----------- src/PostgREST/QueryBuilder.hs | 41 ++++++++++++++++++++++++----------- test/Feature/QuerySpec.hs | 40 ++++++++++++++++++++++++++++++---- test/fixtures/schema.sql | 12 ++++++++++ 4 files changed, 95 insertions(+), 30 deletions(-) diff --git a/src/PostgREST/App.hs b/src/PostgREST/App.hs index 5aeafc711..8101fd480 100644 --- a/src/PostgREST/App.hs +++ b/src/PostgREST/App.hs @@ -187,18 +187,21 @@ app dbStructure conf apiRequest = (ActionInvoke, TargetProc qi, Just (PayloadJSON (UniformObjects payload))) -> do let p = V.head payload + singular = iPreferSingular apiRequest jwtSecret = configJwtSecret conf returnJWT = qiName qi `elem` dbProcsReturningJWT dbStructure - respondToRange $ do - row <- H.query () (callProc qi p topLevelRange shouldCount) - --returnJWT <- H.query qi doesProcReturnJWT - let (tableTotal, queryTotal, body) = fromMaybe (Just 0, 0, emptyArray) row - (status, contentRange) = rangeHeader queryTotal tableTotal - in - return $ responseLBS status [jsonH, contentRange] - (if returnJWT - then "{\"token\":\"" <> cs (tokenJWT jwtSecret body) <> "\"}" - else cs $ encode body) + case readSqlParts of + Left e -> return $ responseLBS status400 [jsonH] $ cs e + Right (q,cq) -> respondToRange $ do + row <- H.query () (callProc qi p q cq topLevelRange shouldCount singular) + --returnJWT <- H.query qi doesProcReturnJWT + let (tableTotal, queryTotal, body) = fromMaybe (Just 0, 0, emptyArray) row + (status, contentRange) = rangeHeader queryTotal tableTotal + in + return $ responseLBS status [jsonH, contentRange] + (if returnJWT + then "{\"token\":\"" <> cs (tokenJWT jwtSecret body) <> "\"}" + else cs $ encode body) (ActionRead, TargetRoot, Nothing) -> do let encodeApi ti = encodeOpenAPI ti host port @@ -323,9 +326,10 @@ addFiltersOrdersRanges apiRequest = foldr1 (liftA2 (.)) [ filters = mapM pRequestFilter flts where action = iAction apiRequest - flts = if action == ActionRead - then iFilters apiRequest - else filter (( '.' `elem` ) . fst) $ iFilters apiRequest -- there can be no filters on the root table whre we are doing insert/update + flts + | action == ActionRead = iFilters apiRequest + | 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 = mapM pRequestOrder $ iOrder apiRequest ranges :: Either ParseError [(Path, NonnegRange)] @@ -347,6 +351,8 @@ buildReadRequest maxRows allRels apiRequest = let target = iTarget apiRequest in case target of (TargetIdent (QualifiedIdentifier s t) ) -> Just (s, t) + (TargetProc (QualifiedIdentifier s p) ) -> Just (s, p) + _ -> Nothing action :: Action diff --git a/src/PostgREST/QueryBuilder.hs b/src/PostgREST/QueryBuilder.hs index 621b5476f..37906635e 100644 --- a/src/PostgREST/QueryBuilder.hs +++ b/src/PostgREST/QueryBuilder.hs @@ -203,29 +203,41 @@ addJoinConditions schema (Node nn@(query, (n, r, a)) forest) = addCond query' con = query'{flt_=con ++ flt_ query'} type ProcResults = (Maybe Int64, Int64, JSON.Value) -callProc :: QualifiedIdentifier -> JSON.Object -> NonnegRange -> Bool -> H.Query () (Maybe ProcResults) -callProc qi params range countTotal = +callProc :: QualifiedIdentifier -> JSON.Object -> SqlQuery -> SqlQuery -> NonnegRange -> Bool -> Bool -> H.Query () (Maybe ProcResults) +callProc qi params selectQuery countQuery _ countTotal isSingle = unicodeStatement sql HE.unit decodeProc True where sql = [qc| - WITH t AS (select * {_callSql}) + WITH {sourceCTEName} AS ({_callSql}) SELECT - {_countExpr} as countTotal, - pg_catalog.count(1) as countResult, - array_to_json( - coalesce(array_agg(row_to_json(r)), '\{}') - )::character varying - FROM (select * from t {limitF range}) r; + {countResultF} AS total_result_set, + pg_catalog.count(t) AS page_total, + case when pg_catalog.count(1) > 1 + then {bodyF} + else ( + select case when ((array_agg(row_to_json(t)))[1]->{_procName}) is not null + then ((array_agg(row_to_json(t)))[1]->{_procName})::character varying + else {bodyF} + end + ) + 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) + _procName = pgFmtLit $ qiName qi _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 - then "(select pg_catalog.count(1) from t)" + then [qc|(select pg_catalog.count(1) from {sourceCTEName})|] else "null::bigint" :: Text decodeProc = HD.maybeRow procRow procRow = (,,) <$> HD.nullableValue HD.int8 <*> HD.value HD.int8 <*> HD.value HD.json + bodyF + | isSingle = asJsonSingleF + | otherwise = asJsonF operators :: [(Text, SqlFragment)] operators = [ @@ -263,10 +275,13 @@ requestToCountQuery _ (DbMutate _) = undefined requestToCountQuery schema (DbRead (Node (Select _ _ conditions _ _, (mainTbl, _, _)) _)) = unwords [ "SELECT pg_catalog.count(1)", - "FROM ", fromQi $ QualifiedIdentifier schema mainTbl, - ("WHERE " <> intercalate " AND " ( map (pgFmtCondition (QualifiedIdentifier schema mainTbl)) localConditions )) `emptyOnNull` localConditions + "FROM ", fromQi qi, + ("WHERE " <> intercalate " AND " ( map (pgFmtCondition qi) localConditions )) `emptyOnNull` localConditions ] where + qi = if mainTbl == sourceCTEName + then QualifiedIdentifier "" mainTbl + else QualifiedIdentifier schema mainTbl fn Filter{value=VText _} = True fn Filter{value=VForeignKey _ _} = False localConditions = filter fn conditions diff --git a/test/Feature/QuerySpec.hs b/test/Feature/QuerySpec.hs index 24ae58671..5f63542a6 100644 --- a/test/Feature/QuerySpec.hs +++ b/test/Feature/QuerySpec.hs @@ -454,6 +454,38 @@ spec = do post "/rpc/getitemrange" [json| { "min": 2, "max": 4 } |] `shouldRespondWith` [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":2},"tasks":[{"id":1}]}]|] + context "a proc that returns an empty rowset" $ it "returns empty json array" $ post "/rpc/test_empty_rowset" [json| {} |] `shouldRespondWith` @@ -462,11 +494,11 @@ spec = do context "a proc that returns plain text" $ do it "returns proper json" $ post "/rpc/sayhello" [json| { "name": "world" } |] `shouldRespondWith` - [json| [{"sayhello":"Hello, world"}] |] + [json|"Hello, world"|] it "can handle unicode" $ post "/rpc/sayhello" [json| { "name": "¥" } |] `shouldRespondWith` - [json| [{"sayhello":"Hello, ¥"}] |] + [json|"Hello, ¥"|] context "improper input" $ do it "rejects unknown content type even if payload is good" $ @@ -502,9 +534,9 @@ spec = do it "executes the proc exactly once per request" $ do post "/rpc/callcounter" [json| {} |] `shouldRespondWith` - [json| [{"callcounter":1}] |] + [json|1|] post "/rpc/callcounter" [json| {} |] `shouldRespondWith` - [json| [{"callcounter":2}] |] + [json|2|] describe "weird requests" $ do it "can query as normal" $ do diff --git a/test/fixtures/schema.sql b/test/fixtures/schema.sql index 17a825de6..bc3973cc9 100755 --- a/test/fixtures/schema.sql +++ b/test/fixtures/schema.sql @@ -1021,6 +1021,18 @@ create table orders ( 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 --