Skip to content

Commit f0da5a5

Browse files
fixes
1 parent d081bec commit f0da5a5

6 files changed

Lines changed: 54 additions & 44 deletions

File tree

packages/bigframes/bigframes/core/compile/substrait/compiler.py

Lines changed: 16 additions & 15 deletions
Original file line numberDiff line numberDiff line change
@@ -30,7 +30,7 @@
3030
import bigframes.operations.generic_ops as generic_ops
3131
import bigframes.operations.numeric_ops as numeric_ops
3232
import bigframes.operations.struct_ops as struct_ops
33-
from bigframes.core import bigframe_node, nodes, rewrite
33+
from bigframes.core import agg_expressions, bigframe_node, nodes, rewrite
3434
from bigframes.core.compile import lowering
3535

3636

@@ -499,12 +499,12 @@ def _compile_window(self, node: nodes.WindowOpNode) -> algebra_pb2.Rel:
499499
bounds_type = (
500500
algebra_pb2.Expression.WindowFunction.BoundsType.BOUNDS_TYPE_RANGE
501501
)
502-
start = node.window_spec.bounds.start
503-
if start is None:
502+
range_start = node.window_spec.bounds.start
503+
if range_start is None:
504504
lower_bound.unbounded.CopyFrom(
505505
algebra_pb2.Expression.WindowFunction.Bound.Unbounded()
506506
)
507-
elif start == pd.Timedelta(0):
507+
elif range_start == pd.Timedelta(0):
508508
lower_bound.current_row.CopyFrom(
509509
algebra_pb2.Expression.WindowFunction.Bound.CurrentRow()
510510
)
@@ -513,12 +513,12 @@ def _compile_window(self, node: nodes.WindowOpNode) -> algebra_pb2.Rel:
513513
"Range window bounds with non-zero offsets are not supported yet"
514514
)
515515

516-
end = node.window_spec.bounds.end
517-
if end is None:
516+
range_end = node.window_spec.bounds.end
517+
if range_end is None:
518518
upper_bound.unbounded.CopyFrom(
519519
algebra_pb2.Expression.WindowFunction.Bound.Unbounded()
520520
)
521-
elif end == pd.Timedelta(0):
521+
elif range_end == pd.Timedelta(0):
522522
upper_bound.current_row.CopyFrom(
523523
algebra_pb2.Expression.WindowFunction.Bound.CurrentRow()
524524
)
@@ -540,6 +540,7 @@ def _compile_window(self, node: nodes.WindowOpNode) -> algebra_pb2.Rel:
540540
# 3. Project each window aggregation expression as a WindowFunction expression
541541
for agg_idx, col_def in enumerate(node.agg_exprs):
542542
agg = col_def.expression
543+
assert isinstance(agg, agg_expressions.Aggregation)
543544
distinct = False
544545

545546
if isinstance(agg.op, agg_ops.SumOp):
@@ -798,10 +799,10 @@ def _compile_scalar_constant(
798799
pb_expr.literal.precision_timestamp.precision = 6
799800
pb_expr.literal.precision_timestamp.value = us
800801
elif isinstance(val, datetime.date):
801-
epoch = datetime.date(1970, 1, 1)
802-
days = (val - epoch).days
802+
date_epoch = datetime.date(1970, 1, 1)
803+
days = (val - date_epoch).days
803804
pb_expr.literal.date = days
804-
elif pd.isna(val):
805+
elif pd.isna(val): # type: ignore[call-overload]
805806
pb_expr.literal.null.varchar.length = 0
806807
else:
807808
pb_expr.literal.string = str(val)
@@ -911,7 +912,7 @@ def _get_expression_dtype(
911912
import bigframes.dtypes as dtypes
912913

913914
if isinstance(expr, ex.ScalarConstantExpression):
914-
if expr.value is None or pd.isna(expr.value):
915+
if expr.value is None or pd.isna(expr.value): # type: ignore[call-overload]
915916
return None
916917
return expr.dtype or dtypes.infer_literal_type(expr.value)
917918
elif isinstance(expr, ex.DerefOp):
@@ -1110,7 +1111,7 @@ def _compile_isin(
11101111
@_compile_op.register(generic_ops.FillNaOp)
11111112
def _compile_fillna_op(
11121113
self,
1113-
op: generic_ops.FillNaOp,
1114+
op: ops.BinaryOp,
11141115
inputs: Sequence[ex.Expression],
11151116
child: nodes.BigFrameNode,
11161117
) -> algebra_pb2.Expression:
@@ -1237,7 +1238,7 @@ def _compile_standard_unaryops(
12371238
@_compile_op.register(numeric_ops.PosOp)
12381239
def _compile_pos_op(
12391240
self,
1240-
op: numeric_ops.PosOp,
1241+
op: ops.UnaryOp,
12411242
inputs: Sequence[ex.Expression],
12421243
child: nodes.BigFrameNode,
12431244
) -> algebra_pb2.Expression:
@@ -1247,7 +1248,7 @@ def _compile_pos_op(
12471248
@_compile_op.register(numeric_ops.NegOp)
12481249
def _compile_neg_op(
12491250
self,
1250-
op: numeric_ops.NegOp,
1251+
op: ops.UnaryOp,
12511252
inputs: Sequence[ex.Expression],
12521253
child: nodes.BigFrameNode,
12531254
) -> algebra_pb2.Expression:
@@ -1270,7 +1271,7 @@ def _compile_neg_op(
12701271
@_compile_op.register(generic_ops.InvertOp)
12711272
def _compile_invert_op(
12721273
self,
1273-
op: generic_ops.InvertOp,
1274+
op: ops.UnaryOp,
12741275
inputs: Sequence[ex.Expression],
12751276
child: nodes.BigFrameNode,
12761277
) -> algebra_pb2.Expression:

packages/bigframes/bigframes/core/rewrite/substrait_agg.py

Lines changed: 12 additions & 22 deletions
Original file line numberDiff line numberDiff line change
@@ -48,7 +48,7 @@ def rewrite_substrait_aggregations(node: nodes.BigFrameNode) -> nodes.BigFrameNo
4848
]
4949

5050
# Collect cast aggregations (bool->float for mean/stddev/var, bool->int for sum)
51-
cast_aggs = []
51+
cast_aggs: list[tuple[int, identifiers.ColumnId, dtypes.Dtype]] = []
5252
for agg_idx, (agg, _) in enumerate(node.aggregations):
5353
if hasattr(agg, "column_references"):
5454
for col_id in agg.column_references:
@@ -92,6 +92,7 @@ def rewrite_substrait_aggregations(node: nodes.BigFrameNode) -> nodes.BigFrameNo
9292
# Rewrite aggregations to use the projected columns
9393
rewritten_aggs = []
9494
for agg_idx, (agg, out_col_id) in enumerate(node.aggregations):
95+
rewritten_agg: agg_expressions.Aggregation
9596
if isinstance(agg.op, (agg_ops.SizeOp, agg_ops.SizeUnaryOp)):
9697
new_col_id = size_agg_to_col_id[agg_idx]
9798
rewritten_agg = agg_expressions.UnaryAggregation(
@@ -107,14 +108,7 @@ def rewrite_substrait_aggregations(node: nodes.BigFrameNode) -> nodes.BigFrameNo
107108
else:
108109
new_exprs.append(expression.deref(col_id.name))
109110

110-
if isinstance(agg, agg_expressions.UnaryAggregation):
111-
rewritten_agg = agg_expressions.UnaryAggregation(
112-
agg.op, new_exprs[0]
113-
)
114-
else:
115-
rewritten_agg = agg_expressions.UnaryAggregation(
116-
agg.op, new_exprs[0]
117-
)
111+
rewritten_agg = agg.replace_args(*new_exprs)
118112
else:
119113
rewritten_agg = agg
120114

@@ -135,9 +129,9 @@ def rewrite_substrait_aggregations(node: nodes.BigFrameNode) -> nodes.BigFrameNo
135129

136130
assignments = []
137131
selection_pairs = []
138-
for idx, (out_id, target_dtype) in enumerate(zip(output_ids, expected_types)):
132+
for idx, (out_id, out_dtype) in enumerate(zip(output_ids, expected_types)):
139133
cast_id = identifiers.ColumnId(f"bf_out_cast_{out_id.name}")
140-
cast_expr = ops.AsTypeOp(to_type=target_dtype).as_expr(
134+
cast_expr = ops.AsTypeOp(to_type=out_dtype).as_expr(
141135
expression.deref(out_id.name)
142136
)
143137
assignments.append((cast_expr, cast_id))
@@ -173,10 +167,11 @@ def rewrite_substrait_windows(node: nodes.BigFrameNode) -> nodes.BigFrameNode:
173167

174168
# Collect size and cast requirements for the window agg expressions
175169
size_aggs = []
176-
cast_aggs = []
170+
cast_aggs: list[tuple[int, identifiers.ColumnId, dtypes.Dtype]] = []
177171

178172
for agg_idx, col_def in enumerate(node.agg_exprs):
179173
agg = col_def.expression
174+
assert isinstance(agg, agg_expressions.Aggregation)
180175
if isinstance(agg.op, (agg_ops.SizeOp, agg_ops.SizeUnaryOp)):
181176
size_aggs.append((agg_idx, agg))
182177
elif hasattr(agg, "column_references"):
@@ -222,7 +217,9 @@ def rewrite_substrait_windows(node: nodes.BigFrameNode) -> nodes.BigFrameNode:
222217
rewritten_agg_exprs = []
223218
for agg_idx, col_def in enumerate(node.agg_exprs):
224219
agg = col_def.expression
220+
assert isinstance(agg, agg_expressions.Aggregation)
225221
out_col_id = col_def.id
222+
rewritten_agg: agg_expressions.Aggregation
226223

227224
if isinstance(agg.op, (agg_ops.SizeOp, agg_ops.SizeUnaryOp)):
228225
new_col_id = size_agg_to_col_id[agg_idx]
@@ -239,14 +236,7 @@ def rewrite_substrait_windows(node: nodes.BigFrameNode) -> nodes.BigFrameNode:
239236
else:
240237
new_exprs.append(expression.deref(col_id.name))
241238

242-
if isinstance(agg, agg_expressions.UnaryAggregation):
243-
rewritten_agg = agg_expressions.UnaryAggregation(
244-
agg.op, new_exprs[0]
245-
)
246-
else:
247-
rewritten_agg = agg_expressions.UnaryAggregation(
248-
agg.op, new_exprs[0]
249-
)
239+
rewritten_agg = agg.replace_args(*new_exprs)
250240
else:
251241
rewritten_agg = agg
252242

@@ -273,10 +263,10 @@ def rewrite_substrait_windows(node: nodes.BigFrameNode) -> nodes.BigFrameNode:
273263

274264
for col_def, field in zip(node.agg_exprs, node.added_fields):
275265
out_id = col_def.id
276-
target_dtype = field.dtype
266+
out_dtype = field.dtype
277267

278268
cast_id = identifiers.ColumnId(f"bf_window_out_cast_{out_id.name}")
279-
cast_expr = ops.AsTypeOp(to_type=target_dtype).as_expr(
269+
cast_expr = ops.AsTypeOp(to_type=out_dtype).as_expr(
280270
expression.deref(out_id.name)
281271
)
282272
assignments.append((cast_expr, cast_id))

packages/bigframes/bigframes/session/substrait_executor.py

Lines changed: 4 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -21,7 +21,7 @@
2121
import bigframes.core.compile.substrait.compiler as substrait_compiler
2222
import bigframes.core.rewrite as rewrite
2323
from bigframes.core import bigframe_node, nodes
24-
from bigframes.session import executor, semi_executor
24+
from bigframes.session import execution_spec, executor, semi_executor
2525

2626
if TYPE_CHECKING:
2727
import pyarrow as pa
@@ -129,9 +129,10 @@ def default_for_engine(cls, engine_name: str) -> SubstraitExecutor:
129129
async def execute(
130130
self,
131131
plan: bigframe_node.BigFrameNode,
132-
ordered: bool,
133-
peek: Optional[int] = None,
132+
execution_spec: execution_spec.ExecutionSpec,
134133
) -> Optional[executor.ExecuteResult]:
134+
ordered = execution_spec.ordered
135+
peek = execution_spec.peek
135136
plan = plan.bottom_up(rewrite.rewrite_slice)
136137
# Only needed for acero technically, datafusion can handle timedeltas
137138
plan = plan.bottom_up(rewrite.rewrite_timedelta_expressions)

packages/bigframes/bigframes/testing/substrait_session.py

Lines changed: 15 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -35,9 +35,18 @@ class SubstraitTestExecutor(bigframes.session.executor.Executor):
3535
def __init__(
3636
self, consumer: bigframes.session.substrait_executor.SubstraitConsumer
3737
):
38-
from bigframes.session.substrait_executor import SubstraitExecutor
38+
from bigframes.core.compile.substrait.compiler import SubstraitCompiler
39+
from bigframes.session.substrait_executor import (
40+
AceroSubstraitConsumer,
41+
SubstraitExecutor,
42+
)
43+
44+
if isinstance(consumer, AceroSubstraitConsumer):
45+
compiler = SubstraitCompiler(duration_type="int", use_precision_types=False)
46+
else:
47+
compiler = SubstraitCompiler(duration_type="int")
3948

40-
self.executor = SubstraitExecutor(consumer)
49+
self.executor = SubstraitExecutor(consumer, compiler)
4150

4251
def execute(
4352
self,
@@ -51,7 +60,10 @@ def execute(
5160

5261
result = asyncio.run(
5362
self.executor.execute(
54-
array_value.node, ordered=True, peek=execution_spec.peek
63+
array_value.node,
64+
bigframes.session.execution_spec.ExecutionSpec(
65+
ordered=True, peek=execution_spec.peek
66+
),
5567
)
5668
)
5769
if result is None:

packages/bigframes/mypy.ini

Lines changed: 6 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -45,5 +45,11 @@ ignore_missing_imports = True
4545
[mypy-anywidget]
4646
ignore_missing_imports = True
4747

48+
[mypy-substrait.*]
49+
ignore_missing_imports = True
50+
51+
[mypy-datafusion.*]
52+
ignore_missing_imports = True
53+
4854
[mypy-bigframes_vendored.*]
4955
ignore_errors = True

packages/bigframes/testing/constraints-3.10.txt

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -18,7 +18,7 @@ pandas==1.5.3
1818
pandas-gbq==0.26.1
1919
polars==1.21.0
2020
substrait==0.29.0
21-
datafusion==45.0.0
21+
datafusion==45.2.0
2222
pyarrow==23.0.1
2323
pydata-google-auth==1.8.2
2424
pyiceberg==0.7.1

0 commit comments

Comments
 (0)