diff --git a/crates/polars-plan/src/plans/optimizer/predicate_pushdown/mod.rs b/crates/polars-plan/src/plans/optimizer/predicate_pushdown/mod.rs index dc0a9521d188..32fa0328fa7e 100644 --- a/crates/polars-plan/src/plans/optimizer/predicate_pushdown/mod.rs +++ b/crates/polars-plan/src/plans/optimizer/predicate_pushdown/mod.rs @@ -196,12 +196,7 @@ impl PredicatePushDown { let new_inputs = inputs .map(|node| { let alp = lp_arena.take(node); - let alp = self.push_down( - alp, - init_hashmap(Some(acc_predicates.len())), - lp_arena, - expr_arena, - )?; + let alp = self.push_down(alp, init_hashmap(None), lp_arena, expr_arena)?; lp_arena.replace(node, alp); Ok(node) }) @@ -581,6 +576,12 @@ impl PredicatePushDown { mut slice, sort_options, } => { + let mut local_predicates = Vec::new(); + if slice.is_some() && !acc_predicates.is_empty() { + local_predicates = acc_predicates.into_values().collect(); + acc_predicates = init_hashmap(None); + } + if let Some((offset, len, None)) = slice && by_column.len() == 1 { @@ -600,7 +601,9 @@ impl PredicatePushDown { slice, sort_options, }; - self.pushdown_and_continue(lp, acc_predicates, lp_arena, expr_arena, true) + let lp = + self.pushdown_and_continue(lp, acc_predicates, lp_arena, expr_arena, true)?; + Ok(self.optional_apply_predicate(lp, local_predicates, lp_arena, expr_arena)) }, lp @ Sink { .. } | lp @ SinkMultiple { .. } => { self.pushdown_and_continue(lp, acc_predicates, lp_arena, expr_arena, false) diff --git a/py-polars/tests/unit/lazyframe/test_predicates.py b/py-polars/tests/unit/lazyframe/test_predicates.py index e5232ec39282..a27d71e7c662 100644 --- a/py-polars/tests/unit/lazyframe/test_predicates.py +++ b/py-polars/tests/unit/lazyframe/test_predicates.py @@ -1293,3 +1293,31 @@ def test_projection_pushed_past_join_26693() -> None: plan = a.filter(pl.col.y > 0).join(b, on="x").group_by("x").agg([]).explain() assert plan.index("simple π") > plan.index("INNER JOIN") + + +def test_predicate_pushdown_sort_slice_26803() -> None: + df = pl.DataFrame({"rank": [1, 2, 3, 4, 5], "score": [10, 4, 6, 2, 8]}) + + for lazy, eager in [ + ( + df.lazy().sort("rank").head(3).filter(pl.col("rank") > 2), + df.sort("rank").head(3).filter(pl.col("rank") > 2), + ), + ( + df.lazy().sort("rank").tail(3).filter(pl.col("rank") < 4), + df.sort("rank").tail(3).filter(pl.col("rank") < 4), + ), + ( + df.lazy().sort("rank", descending=True).head(3).filter(pl.col("rank") < 4), + df.sort("rank", descending=True).head(3).filter(pl.col("rank") < 4), + ), + ( + df.lazy().top_k(3, by="rank").filter(pl.col("rank") < 5), + df.top_k(3, by="rank").filter(pl.col("rank") < 5), + ), + ( + df.lazy().bottom_k(3, by="rank").filter(pl.col("rank") > 1), + df.bottom_k(3, by="rank").filter(pl.col("rank") > 1), + ), + ]: + assert_frame_equal(lazy.collect(), eager, check_row_order=False)