From dd05d9bf6bf21d0a713ca99de584206e6464a4e3 Mon Sep 17 00:00:00 2001 From: Ruifeng Zheng Date: Wed, 9 Sep 2026 02:42:34 +0000 Subject: [PATCH 1/3] [MINOR][CONNECT] Improve Python client plan test coverage --- .../sql/tests/connect/test_connect_plan.py | 194 +++++++++++++++++- 1 file changed, 193 insertions(+), 1 deletion(-) diff --git a/python/pyspark/sql/tests/connect/test_connect_plan.py b/python/pyspark/sql/tests/connect/test_connect_plan.py index d3b660f6ccd86..466b332f14116 100644 --- a/python/pyspark/sql/tests/connect/test_connect_plan.py +++ b/python/pyspark/sql/tests/connect/test_connect_plan.py @@ -85,6 +85,67 @@ def test_simple_project(self): self.assertIsNotNone(plan.root, "Root relation must be set") self.assertIsNotNone(plan.root.read) + def test_select_expr(self): + df = self.connect.readTable(table_name=self.tbl_name) + plan = df.selectExpr("col_a + 1", "col_b AS renamed")._plan.to_proto(self.connect) + self.assertEqual( + [ + expression.expression_string.expression + for expression in plan.root.project.expressions + ], + ["col_a + 1", "col_b AS renamed"], + ) + + def test_aggregate(self): + df = self.connect.readTable(table_name=self.tbl_name) + + plan = df.agg({"value": "sum"})._plan.to_proto(self.connect) + aggregate = plan.root.aggregate + self.assertEqual( + aggregate.group_type, proto.Aggregate.GroupType.GROUP_TYPE_GROUPBY + ) + self.assertEqual(len(aggregate.grouping_expressions), 0) + self.assertEqual( + aggregate.aggregate_expressions[0].unresolved_function.function_name, "sum" + ) + + plan = df.groupBy("key").agg(sum("value"))._plan.to_proto(self.connect) + aggregate = plan.root.aggregate + self.assertEqual( + aggregate.group_type, proto.Aggregate.GroupType.GROUP_TYPE_GROUPBY + ) + self.assertEqual( + aggregate.grouping_expressions[0].unresolved_attribute.unparsed_identifier, "key" + ) + + plan = df.rollup("key").agg(sum("value"))._plan.to_proto(self.connect) + self.assertEqual( + plan.root.aggregate.group_type, proto.Aggregate.GroupType.GROUP_TYPE_ROLLUP + ) + + plan = df.cube("key").agg(sum("value"))._plan.to_proto(self.connect) + self.assertEqual( + plan.root.aggregate.group_type, proto.Aggregate.GroupType.GROUP_TYPE_CUBE + ) + + plan = ( + df.groupingSets([["key"], ["category"]], "key", "category") + .agg(sum("value")) + ._plan.to_proto(self.connect) + ) + aggregate = plan.root.aggregate + self.assertEqual( + aggregate.group_type, proto.Aggregate.GroupType.GROUP_TYPE_GROUPING_SETS + ) + self.assertEqual( + [ + expression.unresolved_attribute.unparsed_identifier + for grouping_set in aggregate.grouping_sets + for expression in grouping_set.grouping_set + ], + ["key", "category"], + ) + def test_bitmap_scalar_functions(self): df = self.connect.readTable(table_name=self.tbl_name) plan = df.select( @@ -121,6 +182,40 @@ def test_join_condition(self): )._plan.to_proto(self.connect) self.assertIsNotNone(plan.root.join.join_condition) + def test_lateral_join(self): + left = self.connect.readTable(table_name=self.tbl_name) + right = self.connect.readTable(table_name=self.tbl_name) + plan = left.lateralJoin(right, left.key == right.key, "left")._plan.to_proto(self.connect) + lateral_join = plan.root.lateral_join + self.assertTrue(lateral_join.HasField("left")) + self.assertTrue(lateral_join.HasField("right")) + self.assertTrue(lateral_join.HasField("join_condition")) + self.assertEqual( + lateral_join.join_type, proto.Join.JoinType.JOIN_TYPE_LEFT_OUTER + ) + + def test_nearest_by_join(self): + left = self.connect.readTable(table_name=self.tbl_name) + right = self.connect.readTable(table_name=self.tbl_name) + plan = left.nearestByJoin( + right, + left.score - right.score, + 3, + "exact", + "distance", + joinType="left", + )._plan.to_proto(self.connect) + nearest_by_join = plan.root.nearest_by_join + self.assertTrue(nearest_by_join.HasField("left")) + self.assertTrue(nearest_by_join.HasField("right")) + self.assertEqual( + nearest_by_join.ranking_expression.unresolved_function.function_name, "-" + ) + self.assertEqual(nearest_by_join.num_results, 3) + self.assertEqual(nearest_by_join.join_type, "left") + self.assertEqual(nearest_by_join.mode, "exact") + self.assertEqual(nearest_by_join.direction, "distance") + def test_crossjoin(self): # SPARK-41227: Test CrossJoin left_input = self.connect.readTable(table_name=self.tbl_name) @@ -154,6 +249,17 @@ def test_zip(self): self.tbl_name, ) + def test_zip_with_index(self): + df = self.connect.readTable(table_name=self.tbl_name) + plan = df.zipWithIndex("row_id")._plan.to_proto(self.connect) + expressions = plan.root.project.expressions + self.assertTrue(expressions[0].HasField("unresolved_star")) + self.assertEqual(expressions[1].alias.name, ["row_id"]) + self.assertEqual( + expressions[1].alias.expr.unresolved_function.function_name, + "distributed_sequence_id", + ) + def test_filter(self): df = self.connect.readTable(table_name=self.tbl_name) plan = df.filter(df.col_name > 3)._plan.to_proto(self.connect) @@ -608,6 +714,12 @@ def test_deduplicate(self): ) self.assertEqual(len(deduplicate_on_subset_columns_plan.root.deduplicate.column_names), 2) + within_watermark_plan = df.dropDuplicatesWithinWatermark(["name"])._plan.to_proto( + self.connect + ) + self.assertTrue(within_watermark_plan.root.deduplicate.within_watermark) + self.assertEqual(within_watermark_plan.root.deduplicate.column_names, ["name"]) + def test_relation_alias(self): df = self.connect.readTable(table_name=self.tbl_name) plan = df.alias("table_alias")._plan.to_proto(self.connect) @@ -721,7 +833,7 @@ def test_union(self): plan1 = df1.union(df2)._plan.to_proto(self.connect) self.assertTrue(plan1.root.set_op.is_all) self.assertEqual(proto.SetOperation.SET_OP_TYPE_UNION, plan1.root.set_op.set_op_type) - plan2 = df1.union(df2)._plan.to_proto(self.connect) + plan2 = df1.unionAll(df2)._plan.to_proto(self.connect) self.assertTrue(plan2.root.set_op.is_all) self.assertEqual(proto.SetOperation.SET_OP_TYPE_UNION, plan2.root.set_op.set_op_type) plan3 = df1.unionByName(df2, True)._plan.to_proto(self.connect) @@ -808,6 +920,19 @@ def test_repartition_by_range(self): ["col_a", "col_b"], ) + def test_repartition_by_id(self): + df = self.connect.readTable(table_name=self.tbl_name) + plan = df.repartitionById(8, "partition_id")._plan.to_proto(self.connect) + repartition = plan.root.repartition_by_expression + self.assertEqual(repartition.num_partitions, 8) + self.assertEqual(len(repartition.partition_exprs), 1) + self.assertEqual( + repartition.partition_exprs[ + 0 + ].direct_shuffle_partition_id.child.unresolved_attribute.unparsed_identifier, + "partition_id", + ) + def test_to(self): # SPARK-41464: test `to` API in Python client. df = self.connect.readTable(table_name=self.tbl_name) @@ -867,6 +992,73 @@ def test_timestamp_nanos_datatype_conversion(self): ), ) + def test_with_columns(self): + df = self.connect.readTable(table_name=self.tbl_name) + + plan = df.withColumn("constant", lit(1))._plan.to_proto(self.connect) + alias = plan.root.with_columns.aliases[0] + self.assertEqual(alias.name, ["constant"]) + self.assertEqual(alias.expr.literal.integer, 1) + + plan = df.withColumns( + {"constant": lit(1), "copied": df.source} + )._plan.to_proto(self.connect) + aliases = plan.root.with_columns.aliases + self.assertEqual([alias.name[0] for alias in aliases], ["constant", "copied"]) + self.assertEqual(aliases[0].expr.literal.integer, 1) + self.assertEqual( + aliases[1].expr.unresolved_attribute.unparsed_identifier, + "source", + ) + + plan = df.withMetadata("source", {"origin": "test"})._plan.to_proto(self.connect) + alias = plan.root.with_columns.aliases[0] + self.assertEqual(alias.name, ["source"]) + self.assertEqual(alias.metadata, '{"origin": "test"}') + + def test_with_columns_renamed(self): + df = self.connect.readTable(table_name=self.tbl_name) + plan = df.withColumnRenamed("old", "new")._plan.to_proto(self.connect) + renames = plan.root.with_columns_renamed.renames + self.assertEqual( + [(rename.col_name, rename.new_col_name) for rename in renames], + [("old", "new")], + ) + + plan = df.withColumnsRenamed({"a": "x", "b": "y"})._plan.to_proto(self.connect) + renames = plan.root.with_columns_renamed.renames + self.assertEqual( + [(rename.col_name, rename.new_col_name) for rename in renames], + [("a", "x"), ("b", "y")], + ) + + def test_with_watermark(self): + df = self.connect.readTable(table_name=self.tbl_name) + plan = df.withWatermark("event_time", "10 minutes")._plan.to_proto(self.connect) + self.assertEqual(plan.root.with_watermark.event_time, "event_time") + self.assertEqual(plan.root.with_watermark.delay_threshold, "10 minutes") + + def test_hint(self): + df = self.connect.readTable(table_name=self.tbl_name) + plan = df.hint("REPARTITION", 8, "key")._plan.to_proto(self.connect) + hint = plan.root.hint + self.assertEqual(hint.name, "REPARTITION") + self.assertEqual(hint.parameters[0].literal.integer, 8) + self.assertEqual(hint.parameters[1].literal.string, "key") + + def test_transpose(self): + df = self.connect.readTable(table_name=self.tbl_name) + plan = df.transpose("key")._plan.to_proto(self.connect) + self.assertEqual( + plan.root.transpose.index_columns[0].unresolved_attribute.unparsed_identifier, + "key", + ) + + def test_to_df(self): + df = self.connect.readTable(table_name=self.tbl_name) + plan = df.toDF("first", "second")._plan.to_proto(self.connect) + self.assertEqual(plan.root.to_df.column_names, ["first", "second"]) + def test_write_operation(self): wo = WriteOperation(self.connect.readTable("name")._plan) wo.mode = "overwrite" From 05fcf23aba277f6635f1fd3a3edf6dc576aef9e1 Mon Sep 17 00:00:00 2001 From: Ruifeng Zheng Date: Wed, 9 Sep 2026 06:14:14 +0000 Subject: [PATCH 2/3] [SPARK-59357][CONNECT] Add AsOfJoin plan coverage --- .../sql/tests/connect/test_connect_plan.py | 30 +++++++++++++++++++ 1 file changed, 30 insertions(+) diff --git a/python/pyspark/sql/tests/connect/test_connect_plan.py b/python/pyspark/sql/tests/connect/test_connect_plan.py index 466b332f14116..85bbe6e724bdc 100644 --- a/python/pyspark/sql/tests/connect/test_connect_plan.py +++ b/python/pyspark/sql/tests/connect/test_connect_plan.py @@ -182,6 +182,36 @@ def test_join_condition(self): )._plan.to_proto(self.connect) self.assertIsNotNone(plan.root.join.join_condition) + def test_as_of_join(self): + left = self.connect.readTable(table_name=self.tbl_name) + right = self.connect.readTable(table_name=self.tbl_name) + plan = left._joinAsOf( + right, + "left_time", + "right_time", + on=["key1", "key2"], + how="left", + tolerance=lit(10), + allowExactMatches=False, + direction="forward", + )._plan.to_proto(self.connect) + as_of_join = plan.root.as_of_join + self.assertTrue(as_of_join.HasField("left")) + self.assertTrue(as_of_join.HasField("right")) + self.assertEqual( + as_of_join.left_as_of.unresolved_attribute.unparsed_identifier, + "left_time", + ) + self.assertEqual( + as_of_join.right_as_of.unresolved_attribute.unparsed_identifier, + "right_time", + ) + self.assertEqual(as_of_join.using_columns, ["key1", "key2"]) + self.assertEqual(as_of_join.join_type, "left") + self.assertEqual(as_of_join.tolerance.literal.integer, 10) + self.assertFalse(as_of_join.allow_exact_matches) + self.assertEqual(as_of_join.direction, "forward") + def test_lateral_join(self): left = self.connect.readTable(table_name=self.tbl_name) right = self.connect.readTable(table_name=self.tbl_name) From 66c25b9201b3b02e7273c15a7b71c35c299968d3 Mon Sep 17 00:00:00 2001 From: Ruifeng Zheng Date: Thu, 10 Sep 2026 06:48:19 +0000 Subject: [PATCH 3/3] [SPARK-59357][CONNECT][TEST] Fix Python plan test formatting --- .../sql/tests/connect/test_connect_plan.py | 30 ++++++------------- 1 file changed, 9 insertions(+), 21 deletions(-) diff --git a/python/pyspark/sql/tests/connect/test_connect_plan.py b/python/pyspark/sql/tests/connect/test_connect_plan.py index 85bbe6e724bdc..685ea3656f5cb 100644 --- a/python/pyspark/sql/tests/connect/test_connect_plan.py +++ b/python/pyspark/sql/tests/connect/test_connect_plan.py @@ -101,9 +101,7 @@ def test_aggregate(self): plan = df.agg({"value": "sum"})._plan.to_proto(self.connect) aggregate = plan.root.aggregate - self.assertEqual( - aggregate.group_type, proto.Aggregate.GroupType.GROUP_TYPE_GROUPBY - ) + self.assertEqual(aggregate.group_type, proto.Aggregate.GroupType.GROUP_TYPE_GROUPBY) self.assertEqual(len(aggregate.grouping_expressions), 0) self.assertEqual( aggregate.aggregate_expressions[0].unresolved_function.function_name, "sum" @@ -111,9 +109,7 @@ def test_aggregate(self): plan = df.groupBy("key").agg(sum("value"))._plan.to_proto(self.connect) aggregate = plan.root.aggregate - self.assertEqual( - aggregate.group_type, proto.Aggregate.GroupType.GROUP_TYPE_GROUPBY - ) + self.assertEqual(aggregate.group_type, proto.Aggregate.GroupType.GROUP_TYPE_GROUPBY) self.assertEqual( aggregate.grouping_expressions[0].unresolved_attribute.unparsed_identifier, "key" ) @@ -124,9 +120,7 @@ def test_aggregate(self): ) plan = df.cube("key").agg(sum("value"))._plan.to_proto(self.connect) - self.assertEqual( - plan.root.aggregate.group_type, proto.Aggregate.GroupType.GROUP_TYPE_CUBE - ) + self.assertEqual(plan.root.aggregate.group_type, proto.Aggregate.GroupType.GROUP_TYPE_CUBE) plan = ( df.groupingSets([["key"], ["category"]], "key", "category") @@ -134,9 +128,7 @@ def test_aggregate(self): ._plan.to_proto(self.connect) ) aggregate = plan.root.aggregate - self.assertEqual( - aggregate.group_type, proto.Aggregate.GroupType.GROUP_TYPE_GROUPING_SETS - ) + self.assertEqual(aggregate.group_type, proto.Aggregate.GroupType.GROUP_TYPE_GROUPING_SETS) self.assertEqual( [ expression.unresolved_attribute.unparsed_identifier @@ -220,9 +212,7 @@ def test_lateral_join(self): self.assertTrue(lateral_join.HasField("left")) self.assertTrue(lateral_join.HasField("right")) self.assertTrue(lateral_join.HasField("join_condition")) - self.assertEqual( - lateral_join.join_type, proto.Join.JoinType.JOIN_TYPE_LEFT_OUTER - ) + self.assertEqual(lateral_join.join_type, proto.Join.JoinType.JOIN_TYPE_LEFT_OUTER) def test_nearest_by_join(self): left = self.connect.readTable(table_name=self.tbl_name) @@ -238,9 +228,7 @@ def test_nearest_by_join(self): nearest_by_join = plan.root.nearest_by_join self.assertTrue(nearest_by_join.HasField("left")) self.assertTrue(nearest_by_join.HasField("right")) - self.assertEqual( - nearest_by_join.ranking_expression.unresolved_function.function_name, "-" - ) + self.assertEqual(nearest_by_join.ranking_expression.unresolved_function.function_name, "-") self.assertEqual(nearest_by_join.num_results, 3) self.assertEqual(nearest_by_join.join_type, "left") self.assertEqual(nearest_by_join.mode, "exact") @@ -1030,9 +1018,9 @@ def test_with_columns(self): self.assertEqual(alias.name, ["constant"]) self.assertEqual(alias.expr.literal.integer, 1) - plan = df.withColumns( - {"constant": lit(1), "copied": df.source} - )._plan.to_proto(self.connect) + plan = df.withColumns({"constant": lit(1), "copied": df.source})._plan.to_proto( + self.connect + ) aliases = plan.root.with_columns.aliases self.assertEqual([alias.name[0] for alias in aliases], ["constant", "copied"]) self.assertEqual(aliases[0].expr.literal.integer, 1)