diff --git a/CHANGELOG.md b/CHANGELOG.md index 4daa473a4..c57a697d3 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -35,6 +35,7 @@ This project adheres to [Semantic Versioning](http://semver.org/). - #3697, #3602, Handle queries on non-existing table gracefully - @taimoorzaeem - #3600, #3926, Improve JWT errors - @taimoorzaeem - #3013, Fix `order=` with POST, PATCH, PUT and DELETE requests - @taimoorzaeem + - #3965, Fix filter on unselected columns in a table-valued function - @taimoorzaeem ### Changed diff --git a/src/PostgREST/Plan.hs b/src/PostgREST/Plan.hs index e5c205ffa..fea81c2b3 100644 --- a/src/PostgREST/Plan.hs +++ b/src/PostgREST/Plan.hs @@ -1019,6 +1019,7 @@ callPlan proc ApiRequest{} paramKeys args readReq = FunctionCall { , funCScalar = funcReturnsScalar proc , funCSetOfScalar = funcReturnsSetOfScalar proc , funCRetCompositeAlias = funcReturnsCompositeAlias proc +, funCFilterFields = getFilterFieldNames readReq , funCReturning = inferColsEmbedNeeds readReq [] } where @@ -1028,6 +1029,22 @@ callPlan proc ApiRequest{} paramKeys args readReq = FunctionCall { | otherwise -> KeyParams $ specifiedParams [prm] prms -> KeyParams $ specifiedParams prms +-- | Get filter fields/column names from read plan +getFilterFieldNames :: ReadPlanTree -> Set FieldName +getFilterFieldNames rpt = S.fromList $ foldr (\rp names -> names <> rpToFieldNames rp) [] rpt + where + rpToFieldNames :: ReadPlan -> [FieldName] + rpToFieldNames = logicTreesToFieldName . ReadPlan.where_ + + logicTreesToFieldName :: [CoercibleLogicTree] -> [FieldName] + logicTreesToFieldName = concatMap coLogicTreeToFieldNames + + coLogicTreeToFieldNames :: CoercibleLogicTree -> [FieldName] + coLogicTreeToFieldNames = \case + CoercibleStmnt (CoercibleFilter{field=CoercibleField{cfName}}) -> [cfName] + CoercibleStmnt (CoercibleFilterNullEmbed _ cfName) -> [cfName] -- needs test coverage + CoercibleExpr _ _ clts -> concatMap coLogicTreeToFieldNames clts + -- | Infers the columns needed for an embed to be successful after a mutation or a function call. inferColsEmbedNeeds :: ReadPlanTree -> [FieldName] -> S.Set FieldName inferColsEmbedNeeds (Node ReadPlan{select} forest) pkCols diff --git a/src/PostgREST/Plan/CallPlan.hs b/src/PostgREST/Plan/CallPlan.hs index 47b06a2c6..2174e82d3 100644 --- a/src/PostgREST/Plan/CallPlan.hs +++ b/src/PostgREST/Plan/CallPlan.hs @@ -25,6 +25,7 @@ data CallPlan = FunctionCall , funCScalar :: Bool , funCSetOfScalar :: Bool , funCRetCompositeAlias :: Bool + , funCFilterFields :: Set FieldName , funCReturning :: Set FieldName } diff --git a/src/PostgREST/Query/QueryBuilder.hs b/src/PostgREST/Query/QueryBuilder.hs index cc108b388..3f7e43244 100644 --- a/src/PostgREST/Query/QueryBuilder.hs +++ b/src/PostgREST/Query/QueryBuilder.hs @@ -170,7 +170,7 @@ mutatePlanToQuery (Delete mainQi logicForest returnings) = whereLogic = if null logicForest then mempty else " WHERE " <> intercalateSnippet " AND " (pgFmtLogicTree mainQi <$> logicForest) callPlanToQuery :: CallPlan -> PgVersion -> SQL.Snippet -callPlanToQuery (FunctionCall qi params arguments returnsScalar returnsSetOfScalar returnsCompositeAlias returnings) pgVer = +callPlanToQuery (FunctionCall qi params arguments returnsScalar returnsSetOfScalar returnsCompositeAlias filterFields returnings) pgVer = "SELECT " <> (if returnsScalar || returnsSetOfScalar then "pgrst_call.pgrst_scalar" else returnedColumns) <> " " <> fromCall where @@ -211,10 +211,16 @@ callPlanToQuery (FunctionCall qi params arguments returnsScalar returnsSetOfScal -- We could fallback to providing this NULL value in those cases. encodeArg Nothing = "NULL" + -- the columns here would be the returnings + the columns that would later + -- be used by a where clause filter, if they intersect, we remove the duplicates + -- and if * is returned then no need to explicitly add filter columns returnedColumns :: SQL.Snippet - returnedColumns - | null returnings = "*" - | otherwise = intercalateSnippet ", " (pgFmtColumn (QualifiedIdentifier mempty "pgrst_call") <$> S.toList returnings) + returnedColumns = case S.toList returnings of + [] -> "*" + ["*"] -> pgFmtColumn (QualifiedIdentifier mempty "pgrst_call") "*" + _ -> intercalateSnippet ", " (pgFmtColumn (QualifiedIdentifier mempty "pgrst_call") <$> returnedColumns') + where + returnedColumns' = S.toList $ returnings <> filterFields -- | SQL query meant for COUNTing the root node of the Tree. -- It only takes WHERE into account and doesn't include LIMIT/OFFSET because it would reduce the COUNT. diff --git a/test/spec/Feature/Query/RpcSpec.hs b/test/spec/Feature/Query/RpcSpec.hs index 50375f9ed..886791110 100644 --- a/test/spec/Feature/Query/RpcSpec.hs +++ b/test/spec/Feature/Query/RpcSpec.hs @@ -1422,3 +1422,31 @@ spec = , matchHeaders = [ "Content-Length" <:> "105" , matchContentTypeJson ] } + + context "test table valued function with filter" $ do + it "works with filter on unselected columns" $ + request methodGet "/rpc/getallprojects?select=id,client_id&name=like.OSX" + [] "" + `shouldRespondWith` + [json| [{"id":4,"client_id":2}] |] + { matchStatus = 200 + , matchHeaders = [matchContentTypeJson] + } + + it "works with filter on unselected columns with null embed" $ + request methodGet "/rpc/getallprojects?select=id,clients(id)&clients.name=not.is.null" + [] "" + `shouldRespondWith` + [json| [{"id":1,"clients":{"id": 1}}, {"id":2,"clients":{"id": 1}}, {"id":3,"clients":{"id": 2}}, {"id":4,"clients":{"id": 2}}, {"id":5,"clients":null}] |] + { matchStatus = 200 + , matchHeaders = [matchContentTypeJson] + } + + it "works with logical filter on unselected columns" $ + request methodGet "/rpc/getallprojects?select=id,client_id&or=(name.like.OSX,name.like.IOS)" + [] "" + `shouldRespondWith` + [json| [{"id":3,"client_id":2}, {"id":4,"client_id":2}] |] + { matchStatus = 200 + , matchHeaders = [matchContentTypeJson] + }