duckpd 0.0.2__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- duckpd/__init__.py +23 -0
- duckpd/_compiler.py +464 -0
- duckpd/_executor.py +153 -0
- duckpd/_logical.py +387 -0
- duckpd/_merging.py +235 -0
- duckpd/_metadata.py +212 -0
- duckpd/_quoting.py +6 -0
- duckpd/_reductions.py +146 -0
- duckpd/_typing.py +56 -0
- duckpd/accessors.py +122 -0
- duckpd/errors.py +29 -0
- duckpd/frame.py +566 -0
- duckpd/groupby.py +206 -0
- duckpd/io.py +60 -0
- duckpd/py.typed +1 -0
- duckpd/series.py +264 -0
- duckpd/session.py +254 -0
- duckpd-0.0.2.dist-info/METADATA +139 -0
- duckpd-0.0.2.dist-info/RECORD +21 -0
- duckpd-0.0.2.dist-info/WHEEL +4 -0
- duckpd-0.0.2.dist-info/licenses/LICENSE +21 -0
duckpd/__init__.py
ADDED
|
@@ -0,0 +1,23 @@
|
|
|
1
|
+
"""Lazy pandas-shaped DataFrames powered by DuckDB."""
|
|
2
|
+
|
|
3
|
+
from importlib.metadata import version
|
|
4
|
+
|
|
5
|
+
from duckpd.frame import DataFrame
|
|
6
|
+
from duckpd.groupby import DataFrameGroupBy
|
|
7
|
+
from duckpd.io import from_arrow, from_pandas, read_parquet
|
|
8
|
+
from duckpd.series import Series
|
|
9
|
+
from duckpd.session import Session, connect
|
|
10
|
+
|
|
11
|
+
__version__ = version("duckpd")
|
|
12
|
+
|
|
13
|
+
__all__ = [
|
|
14
|
+
"DataFrame",
|
|
15
|
+
"DataFrameGroupBy",
|
|
16
|
+
"Series",
|
|
17
|
+
"Session",
|
|
18
|
+
"__version__",
|
|
19
|
+
"connect",
|
|
20
|
+
"from_arrow",
|
|
21
|
+
"from_pandas",
|
|
22
|
+
"read_parquet",
|
|
23
|
+
]
|
duckpd/_compiler.py
ADDED
|
@@ -0,0 +1,464 @@
|
|
|
1
|
+
"""Compilation of DuckPD logical plans into DuckDB relations."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
from dataclasses import dataclass
|
|
6
|
+
from typing import TYPE_CHECKING
|
|
7
|
+
|
|
8
|
+
import duckdb
|
|
9
|
+
import pandas as pd
|
|
10
|
+
import pyarrow as pa
|
|
11
|
+
|
|
12
|
+
from duckpd._logical import (
|
|
13
|
+
AggregateExpression,
|
|
14
|
+
AggregateOperator,
|
|
15
|
+
AggregatePlan,
|
|
16
|
+
ArrowSource,
|
|
17
|
+
BinaryOperator,
|
|
18
|
+
Column,
|
|
19
|
+
ColumnId,
|
|
20
|
+
ColumnRef,
|
|
21
|
+
Expression,
|
|
22
|
+
FilterPlan,
|
|
23
|
+
FunctionCall,
|
|
24
|
+
JoinPlan,
|
|
25
|
+
JoinType,
|
|
26
|
+
LiteralValue,
|
|
27
|
+
LogicalPlan,
|
|
28
|
+
NullPlacement,
|
|
29
|
+
PandasSource,
|
|
30
|
+
ParquetSource,
|
|
31
|
+
ProjectPlan,
|
|
32
|
+
ScanPlan,
|
|
33
|
+
SortDirection,
|
|
34
|
+
SortKey,
|
|
35
|
+
SortPlan,
|
|
36
|
+
SqlSource,
|
|
37
|
+
TableSource,
|
|
38
|
+
UnaryExpression,
|
|
39
|
+
UnaryOperator,
|
|
40
|
+
)
|
|
41
|
+
from duckpd._quoting import quote_identifier
|
|
42
|
+
|
|
43
|
+
if TYPE_CHECKING:
|
|
44
|
+
from duckpd.session import Session
|
|
45
|
+
|
|
46
|
+
|
|
47
|
+
@dataclass(frozen=True)
|
|
48
|
+
class CompiledFrame:
|
|
49
|
+
"""A DuckDB relation and its logical-to-physical column bindings."""
|
|
50
|
+
|
|
51
|
+
relation: duckdb.DuckDBPyRelation
|
|
52
|
+
bindings: dict[ColumnId, str]
|
|
53
|
+
|
|
54
|
+
|
|
55
|
+
class DuckDBCompiler:
|
|
56
|
+
"""Compile typed DuckPD plans without triggering query output."""
|
|
57
|
+
|
|
58
|
+
def __init__(self, session: Session) -> None:
|
|
59
|
+
self._session = session
|
|
60
|
+
|
|
61
|
+
def inspect_source(
|
|
62
|
+
self,
|
|
63
|
+
source: ArrowSource | PandasSource | ParquetSource | SqlSource | TableSource,
|
|
64
|
+
) -> tuple[Column, ...]:
|
|
65
|
+
relation = self._relation_for_source(source)
|
|
66
|
+
labels = relation.columns
|
|
67
|
+
if len(labels) != len(set(labels)):
|
|
68
|
+
msg = "DuckPD does not yet support duplicate column labels"
|
|
69
|
+
raise ValueError(msg)
|
|
70
|
+
return tuple(
|
|
71
|
+
Column(ColumnId.create(), label, str(duckdb_type))
|
|
72
|
+
for label, duckdb_type in zip(labels, relation.types, strict=True)
|
|
73
|
+
)
|
|
74
|
+
|
|
75
|
+
def project_visible(
|
|
76
|
+
self, compiled: CompiledFrame, plan: LogicalPlan
|
|
77
|
+
) -> CompiledFrame:
|
|
78
|
+
"""Drop hidden metadata columns at non-pandas output boundaries."""
|
|
79
|
+
visible = plan.metadata.visible_columns
|
|
80
|
+
expressions = tuple(
|
|
81
|
+
duckdb.SQLExpression(quote_identifier(compiled.bindings[column.id])).alias(
|
|
82
|
+
column.label
|
|
83
|
+
)
|
|
84
|
+
for column in visible
|
|
85
|
+
)
|
|
86
|
+
relation = compiled.relation.project(*expressions)
|
|
87
|
+
return CompiledFrame(
|
|
88
|
+
relation,
|
|
89
|
+
{column.id: column.label for column in visible},
|
|
90
|
+
)
|
|
91
|
+
|
|
92
|
+
def compile(self, plan: LogicalPlan) -> CompiledFrame:
|
|
93
|
+
self._session._ensure_open()
|
|
94
|
+
if isinstance(plan, ScanPlan):
|
|
95
|
+
relation = self._relation_for_source(plan.source)
|
|
96
|
+
bindings = {column.id: column.label for column in plan.columns}
|
|
97
|
+
return CompiledFrame(relation, bindings)
|
|
98
|
+
|
|
99
|
+
if isinstance(plan, JoinPlan):
|
|
100
|
+
return self._compile_join(plan)
|
|
101
|
+
|
|
102
|
+
compiled_input = self.compile(plan.input)
|
|
103
|
+
if isinstance(plan, FilterPlan):
|
|
104
|
+
predicate = self.compile_expression(plan.predicate, compiled_input.bindings)
|
|
105
|
+
return CompiledFrame(
|
|
106
|
+
compiled_input.relation.filter(predicate), compiled_input.bindings
|
|
107
|
+
)
|
|
108
|
+
if isinstance(plan, ProjectPlan):
|
|
109
|
+
expressions = tuple(
|
|
110
|
+
self.compile_expression(
|
|
111
|
+
projection.expression, compiled_input.bindings
|
|
112
|
+
).alias(projection.column.label)
|
|
113
|
+
for projection in plan.projections
|
|
114
|
+
)
|
|
115
|
+
relation = compiled_input.relation.project(*expressions)
|
|
116
|
+
bindings = {
|
|
117
|
+
projection.column.id: projection.column.label
|
|
118
|
+
for projection in plan.projections
|
|
119
|
+
}
|
|
120
|
+
return CompiledFrame(relation, bindings)
|
|
121
|
+
if isinstance(plan, AggregatePlan):
|
|
122
|
+
input_rel = compiled_input.relation
|
|
123
|
+
if plan.keys and plan.dropna:
|
|
124
|
+
# filter out rows where any group key is null
|
|
125
|
+
for key_id in plan.keys:
|
|
126
|
+
key_label = quote_identifier(compiled_input.bindings[key_id])
|
|
127
|
+
input_rel = input_rel.filter(
|
|
128
|
+
duckdb.SQLExpression(f"{key_label} IS NOT NULL")
|
|
129
|
+
)
|
|
130
|
+
|
|
131
|
+
expressions = [
|
|
132
|
+
self._compile_aggregate(aggregate, compiled_input.bindings).alias(
|
|
133
|
+
aggregate.column.label
|
|
134
|
+
)
|
|
135
|
+
for aggregate in plan.aggregates
|
|
136
|
+
]
|
|
137
|
+
if plan.keys:
|
|
138
|
+
key_labels = [
|
|
139
|
+
quote_identifier(compiled_input.bindings[key_id])
|
|
140
|
+
for key_id in plan.keys
|
|
141
|
+
]
|
|
142
|
+
groups_spec = ", ".join(key_labels)
|
|
143
|
+
relation = input_rel.aggregate(expressions, groups_spec)
|
|
144
|
+
if plan.sort:
|
|
145
|
+
sort_keys = [
|
|
146
|
+
duckdb.SQLExpression(k).asc().nulls_last() for k in key_labels
|
|
147
|
+
]
|
|
148
|
+
relation = relation.sort(*sort_keys)
|
|
149
|
+
else:
|
|
150
|
+
relation = input_rel.aggregate(expressions)
|
|
151
|
+
|
|
152
|
+
return CompiledFrame(
|
|
153
|
+
relation,
|
|
154
|
+
{
|
|
155
|
+
aggregate.column.id: aggregate.column.label
|
|
156
|
+
for aggregate in plan.aggregates
|
|
157
|
+
},
|
|
158
|
+
)
|
|
159
|
+
if isinstance(plan, SortPlan):
|
|
160
|
+
keys = tuple(
|
|
161
|
+
self._compile_sort_key(key, compiled_input.bindings)
|
|
162
|
+
for key in plan.keys
|
|
163
|
+
)
|
|
164
|
+
return CompiledFrame(
|
|
165
|
+
compiled_input.relation.sort(*keys), compiled_input.bindings
|
|
166
|
+
)
|
|
167
|
+
relation = compiled_input.relation.limit(plan.count, offset=plan.offset)
|
|
168
|
+
return CompiledFrame(relation, compiled_input.bindings)
|
|
169
|
+
|
|
170
|
+
def compile_expression(
|
|
171
|
+
self, expression: Expression, bindings: dict[ColumnId, str]
|
|
172
|
+
) -> duckdb.Expression:
|
|
173
|
+
"""Compile a typed expression against physical relation bindings."""
|
|
174
|
+
if isinstance(expression, ColumnRef):
|
|
175
|
+
try:
|
|
176
|
+
label = bindings[expression.column_id]
|
|
177
|
+
except KeyError as error:
|
|
178
|
+
msg = (
|
|
179
|
+
f"Column {expression.column_id.value} is not available in this plan"
|
|
180
|
+
)
|
|
181
|
+
raise KeyError(msg) from error
|
|
182
|
+
return duckdb.SQLExpression(quote_identifier(label))
|
|
183
|
+
if isinstance(expression, LiteralValue):
|
|
184
|
+
return duckdb.ConstantExpression(expression.value)
|
|
185
|
+
if isinstance(expression, UnaryExpression):
|
|
186
|
+
operand = self.compile_expression(expression.operand, bindings)
|
|
187
|
+
if expression.operator is UnaryOperator.INVERT:
|
|
188
|
+
return ~operand
|
|
189
|
+
if expression.operator is UnaryOperator.NEGATE:
|
|
190
|
+
return -operand
|
|
191
|
+
if expression.operator is UnaryOperator.POSITIVE:
|
|
192
|
+
return operand
|
|
193
|
+
raise AssertionError(f"Unknown unary operator: {expression.operator}")
|
|
194
|
+
if isinstance(expression, FunctionCall):
|
|
195
|
+
compiled_args = [
|
|
196
|
+
self.compile_expression(arg, bindings) for arg in expression.arguments
|
|
197
|
+
]
|
|
198
|
+
if expression.name.lower() == "coalesce" and len(compiled_args) == 2:
|
|
199
|
+
return duckdb.CaseExpression(
|
|
200
|
+
compiled_args[0].isnull(), compiled_args[1]
|
|
201
|
+
).otherwise(compiled_args[0])
|
|
202
|
+
return duckdb.FunctionExpression(expression.name, *compiled_args)
|
|
203
|
+
|
|
204
|
+
left = self.compile_expression(expression.left, bindings)
|
|
205
|
+
right = self.compile_expression(expression.right, bindings)
|
|
206
|
+
operator = expression.operator
|
|
207
|
+
if operator is BinaryOperator.ADD:
|
|
208
|
+
return left + right
|
|
209
|
+
if operator is BinaryOperator.SUBTRACT:
|
|
210
|
+
return left - right
|
|
211
|
+
if operator is BinaryOperator.MULTIPLY:
|
|
212
|
+
return left * right
|
|
213
|
+
if operator is BinaryOperator.TRUE_DIVIDE:
|
|
214
|
+
return left / right
|
|
215
|
+
if operator is BinaryOperator.MODULO:
|
|
216
|
+
return left % right
|
|
217
|
+
if operator is BinaryOperator.EQUAL:
|
|
218
|
+
return left == right
|
|
219
|
+
if operator is BinaryOperator.NOT_EQUAL:
|
|
220
|
+
return left != right
|
|
221
|
+
if operator is BinaryOperator.LESS_THAN:
|
|
222
|
+
return left < right
|
|
223
|
+
if operator is BinaryOperator.LESS_EQUAL:
|
|
224
|
+
return left <= right
|
|
225
|
+
if operator is BinaryOperator.GREATER_THAN:
|
|
226
|
+
return left > right
|
|
227
|
+
if operator is BinaryOperator.GREATER_EQUAL:
|
|
228
|
+
return left >= right
|
|
229
|
+
if operator is BinaryOperator.AND:
|
|
230
|
+
return left & right
|
|
231
|
+
if operator is BinaryOperator.OR:
|
|
232
|
+
return left | right
|
|
233
|
+
raise AssertionError(f"Unknown binary operator: {operator}")
|
|
234
|
+
|
|
235
|
+
def _relation_for_source(
|
|
236
|
+
self,
|
|
237
|
+
source: ArrowSource | PandasSource | ParquetSource | SqlSource | TableSource,
|
|
238
|
+
) -> duckdb.DuckDBPyRelation:
|
|
239
|
+
self._session._ensure_open()
|
|
240
|
+
if isinstance(source, PandasSource):
|
|
241
|
+
value = self._session._get_registered_source(source.key)
|
|
242
|
+
if not isinstance(value, pd.DataFrame):
|
|
243
|
+
msg = f"Registered source {source.key!r} is not a pandas DataFrame"
|
|
244
|
+
raise TypeError(msg)
|
|
245
|
+
return self._session._connection.from_df(value)
|
|
246
|
+
if isinstance(source, ArrowSource):
|
|
247
|
+
value = self._session._get_registered_source(source.key)
|
|
248
|
+
if not isinstance(value, (pa.Table, pa.RecordBatch)):
|
|
249
|
+
msg = f"Registered source {source.key!r} is not an Arrow table or batch"
|
|
250
|
+
raise TypeError(msg)
|
|
251
|
+
return self._session._connection.from_arrow(value)
|
|
252
|
+
if isinstance(source, TableSource):
|
|
253
|
+
return self._session._connection.table(source.name)
|
|
254
|
+
if isinstance(source, SqlSource):
|
|
255
|
+
return self._session._connection.sql(source.query)
|
|
256
|
+
|
|
257
|
+
paths: str | list[str] = (
|
|
258
|
+
source.paths[0] if len(source.paths) == 1 else list(source.paths)
|
|
259
|
+
)
|
|
260
|
+
return self._session._connection.read_parquet(
|
|
261
|
+
paths,
|
|
262
|
+
hive_partitioning=source.hive_partitioning,
|
|
263
|
+
union_by_name=source.union_by_name,
|
|
264
|
+
)
|
|
265
|
+
|
|
266
|
+
def _compile_sort_key(
|
|
267
|
+
self, key: SortKey, bindings: dict[ColumnId, str]
|
|
268
|
+
) -> duckdb.Expression:
|
|
269
|
+
result = self.compile_expression(key.expression, bindings)
|
|
270
|
+
result = (
|
|
271
|
+
result.asc() if key.direction is SortDirection.ASCENDING else result.desc()
|
|
272
|
+
)
|
|
273
|
+
return (
|
|
274
|
+
result.nulls_first()
|
|
275
|
+
if key.null_placement is NullPlacement.FIRST
|
|
276
|
+
else result.nulls_last()
|
|
277
|
+
)
|
|
278
|
+
|
|
279
|
+
def _compile_aggregate(
|
|
280
|
+
self,
|
|
281
|
+
aggregate: AggregateExpression,
|
|
282
|
+
bindings: dict[ColumnId, str],
|
|
283
|
+
) -> duckdb.Expression:
|
|
284
|
+
if aggregate.operator is AggregateOperator.SIZE:
|
|
285
|
+
return duckdb.SQLExpression("count(*)")
|
|
286
|
+
if aggregate.expression is None:
|
|
287
|
+
raise AssertionError("Only size aggregates may omit an expression")
|
|
288
|
+
|
|
289
|
+
# Identity column pass-through for group keys in projection
|
|
290
|
+
if aggregate.operator is None:
|
|
291
|
+
return self.compile_expression(aggregate.expression, bindings)
|
|
292
|
+
|
|
293
|
+
operand = self.compile_expression(aggregate.expression, bindings)
|
|
294
|
+
non_null_count = duckdb.FunctionExpression("count", operand)
|
|
295
|
+
if aggregate.operator is AggregateOperator.COUNT:
|
|
296
|
+
return non_null_count
|
|
297
|
+
|
|
298
|
+
function = {
|
|
299
|
+
AggregateOperator.SUM: "sum",
|
|
300
|
+
AggregateOperator.MEAN: "avg",
|
|
301
|
+
AggregateOperator.MIN: "min",
|
|
302
|
+
AggregateOperator.MAX: "max",
|
|
303
|
+
}[aggregate.operator]
|
|
304
|
+
aggregate_operand = operand
|
|
305
|
+
if aggregate.input_duckdb_type == "BOOLEAN" and function in {"sum", "avg"}:
|
|
306
|
+
aggregate_operand = operand.cast("BIGINT")
|
|
307
|
+
value = duckdb.FunctionExpression(function, aggregate_operand)
|
|
308
|
+
if aggregate.operator is AggregateOperator.SUM:
|
|
309
|
+
value = duckdb.CaseExpression(
|
|
310
|
+
non_null_count == duckdb.ConstantExpression(0),
|
|
311
|
+
duckdb.ConstantExpression(0),
|
|
312
|
+
).otherwise(value)
|
|
313
|
+
if aggregate.input_duckdb_type == "BOOLEAN" or (
|
|
314
|
+
aggregate.input_duckdb_type is not None
|
|
315
|
+
and aggregate.input_duckdb_type
|
|
316
|
+
in {"TINYINT", "SMALLINT", "INTEGER", "BIGINT"}
|
|
317
|
+
):
|
|
318
|
+
value = value.cast("BIGINT")
|
|
319
|
+
elif aggregate.input_duckdb_type in {
|
|
320
|
+
"UTINYINT",
|
|
321
|
+
"USMALLINT",
|
|
322
|
+
"UINTEGER",
|
|
323
|
+
"UBIGINT",
|
|
324
|
+
}:
|
|
325
|
+
value = value.cast("UBIGINT")
|
|
326
|
+
|
|
327
|
+
invalid = non_null_count < duckdb.ConstantExpression(aggregate.min_count)
|
|
328
|
+
if not aggregate.skipna:
|
|
329
|
+
row_count = duckdb.SQLExpression("count(*)")
|
|
330
|
+
invalid = invalid | (non_null_count < row_count)
|
|
331
|
+
return duckdb.CaseExpression(
|
|
332
|
+
invalid,
|
|
333
|
+
duckdb.ConstantExpression(None),
|
|
334
|
+
).otherwise(value)
|
|
335
|
+
|
|
336
|
+
def _compile_join(self, plan: JoinPlan) -> CompiledFrame:
|
|
337
|
+
left_compiled = self.compile(plan.left)
|
|
338
|
+
right_compiled = self.compile(plan.right)
|
|
339
|
+
|
|
340
|
+
lhs_alias = "lhs"
|
|
341
|
+
rhs_alias = "rhs"
|
|
342
|
+
|
|
343
|
+
# Explicitly project each side with unique physical column names
|
|
344
|
+
left_proj: list[duckdb.Expression] = []
|
|
345
|
+
left_temp_bindings: dict[ColumnId, str] = {}
|
|
346
|
+
for col in plan.left.columns:
|
|
347
|
+
temp_name = f"l_{col.id.value.hex[:8]}"
|
|
348
|
+
left_proj.append(
|
|
349
|
+
duckdb.SQLExpression(
|
|
350
|
+
quote_identifier(left_compiled.bindings[col.id])
|
|
351
|
+
).alias(temp_name)
|
|
352
|
+
)
|
|
353
|
+
left_temp_bindings[col.id] = temp_name
|
|
354
|
+
|
|
355
|
+
right_proj: list[duckdb.Expression] = []
|
|
356
|
+
right_temp_bindings: dict[ColumnId, str] = {}
|
|
357
|
+
for col in plan.right.columns:
|
|
358
|
+
temp_name = f"r_{col.id.value.hex[:8]}"
|
|
359
|
+
right_proj.append(
|
|
360
|
+
duckdb.SQLExpression(
|
|
361
|
+
quote_identifier(right_compiled.bindings[col.id])
|
|
362
|
+
).alias(temp_name)
|
|
363
|
+
)
|
|
364
|
+
right_temp_bindings[col.id] = temp_name
|
|
365
|
+
|
|
366
|
+
lhs_rel = left_compiled.relation.project(*left_proj).set_alias(lhs_alias)
|
|
367
|
+
rhs_rel = right_compiled.relation.project(*right_proj).set_alias(rhs_alias)
|
|
368
|
+
|
|
369
|
+
if plan.how is JoinType.CROSS:
|
|
370
|
+
join_rel = lhs_rel.cross(rhs_rel)
|
|
371
|
+
elif plan.how is JoinType.INNER:
|
|
372
|
+
join_cond_parts = [
|
|
373
|
+
(
|
|
374
|
+
f"{lhs_alias}.{quote_identifier(left_temp_bindings[l_id])} "
|
|
375
|
+
f"IS NOT DISTINCT FROM "
|
|
376
|
+
f"{rhs_alias}.{quote_identifier(right_temp_bindings[r_id])}"
|
|
377
|
+
)
|
|
378
|
+
for l_id, r_id in zip(plan.left_keys, plan.right_keys, strict=True)
|
|
379
|
+
]
|
|
380
|
+
join_rel = lhs_rel.join(rhs_rel, " AND ".join(join_cond_parts), how="inner")
|
|
381
|
+
elif plan.how is JoinType.LEFT:
|
|
382
|
+
join_cond_parts = [
|
|
383
|
+
(
|
|
384
|
+
f"{lhs_alias}.{quote_identifier(left_temp_bindings[l_id])} "
|
|
385
|
+
f"IS NOT DISTINCT FROM "
|
|
386
|
+
f"{rhs_alias}.{quote_identifier(right_temp_bindings[r_id])}"
|
|
387
|
+
)
|
|
388
|
+
for l_id, r_id in zip(plan.left_keys, plan.right_keys, strict=True)
|
|
389
|
+
]
|
|
390
|
+
join_rel = lhs_rel.join(rhs_rel, " AND ".join(join_cond_parts), how="left")
|
|
391
|
+
elif plan.how is JoinType.RIGHT:
|
|
392
|
+
join_cond_parts = [
|
|
393
|
+
(
|
|
394
|
+
f"{lhs_alias}.{quote_identifier(left_temp_bindings[l_id])} "
|
|
395
|
+
f"IS NOT DISTINCT FROM "
|
|
396
|
+
f"{rhs_alias}.{quote_identifier(right_temp_bindings[r_id])}"
|
|
397
|
+
)
|
|
398
|
+
for l_id, r_id in zip(plan.left_keys, plan.right_keys, strict=True)
|
|
399
|
+
]
|
|
400
|
+
join_rel = lhs_rel.join(rhs_rel, " AND ".join(join_cond_parts), how="right")
|
|
401
|
+
elif plan.how is JoinType.OUTER:
|
|
402
|
+
join_cond_parts = [
|
|
403
|
+
(
|
|
404
|
+
f"{lhs_alias}.{quote_identifier(left_temp_bindings[l_id])} "
|
|
405
|
+
f"IS NOT DISTINCT FROM "
|
|
406
|
+
f"{rhs_alias}.{quote_identifier(right_temp_bindings[r_id])}"
|
|
407
|
+
)
|
|
408
|
+
for l_id, r_id in zip(plan.left_keys, plan.right_keys, strict=True)
|
|
409
|
+
]
|
|
410
|
+
join_rel = lhs_rel.join(rhs_rel, " AND ".join(join_cond_parts), how="outer")
|
|
411
|
+
else:
|
|
412
|
+
raise AssertionError(f"Unknown JoinType: {plan.how}")
|
|
413
|
+
|
|
414
|
+
# Output projection to match plan.metadata.columns
|
|
415
|
+
final_proj: list[duckdb.Expression] = []
|
|
416
|
+
final_bindings: dict[ColumnId, str] = {}
|
|
417
|
+
for col in plan.metadata.columns:
|
|
418
|
+
l_bind = left_temp_bindings.get(col.id)
|
|
419
|
+
r_bind = None
|
|
420
|
+
if col.id in plan.left_keys:
|
|
421
|
+
idx = plan.left_keys.index(col.id)
|
|
422
|
+
r_key_id = plan.right_keys[idx]
|
|
423
|
+
r_bind = right_temp_bindings.get(r_key_id)
|
|
424
|
+
|
|
425
|
+
if (
|
|
426
|
+
l_bind is not None
|
|
427
|
+
and r_bind is not None
|
|
428
|
+
and plan.how in {JoinType.RIGHT, JoinType.OUTER}
|
|
429
|
+
):
|
|
430
|
+
# Coalesce left and right key so right-only rows retain value
|
|
431
|
+
l_col = f"{lhs_alias}.{quote_identifier(l_bind)}"
|
|
432
|
+
r_col = f"{rhs_alias}.{quote_identifier(r_bind)}"
|
|
433
|
+
source_col = f"COALESCE({l_col}, {r_col})"
|
|
434
|
+
elif l_bind is not None:
|
|
435
|
+
source_col = f"{lhs_alias}.{quote_identifier(l_bind)}"
|
|
436
|
+
elif col.id in right_temp_bindings:
|
|
437
|
+
source_col = (
|
|
438
|
+
f"{rhs_alias}.{quote_identifier(right_temp_bindings[col.id])}"
|
|
439
|
+
)
|
|
440
|
+
else:
|
|
441
|
+
# Could happen if a key column was synthesized
|
|
442
|
+
raise AssertionError(
|
|
443
|
+
f"Column {col.id} not found in join input bindings"
|
|
444
|
+
)
|
|
445
|
+
|
|
446
|
+
final_proj.append(duckdb.SQLExpression(source_col).alias(col.label))
|
|
447
|
+
final_bindings[col.id] = col.label
|
|
448
|
+
|
|
449
|
+
result_rel = join_rel.project(*final_proj)
|
|
450
|
+
|
|
451
|
+
if plan.sort and plan.metadata.ordering.keys:
|
|
452
|
+
sort_keys = [
|
|
453
|
+
duckdb.SQLExpression(quote_identifier(final_bindings[k.column_id]))
|
|
454
|
+
.asc()
|
|
455
|
+
.nulls_last()
|
|
456
|
+
if k.direction is SortDirection.ASCENDING
|
|
457
|
+
else duckdb.SQLExpression(quote_identifier(final_bindings[k.column_id]))
|
|
458
|
+
.desc()
|
|
459
|
+
.nulls_last()
|
|
460
|
+
for k in plan.metadata.ordering.keys
|
|
461
|
+
]
|
|
462
|
+
result_rel = result_rel.sort(*sort_keys)
|
|
463
|
+
|
|
464
|
+
return CompiledFrame(result_rel, final_bindings)
|
duckpd/_executor.py
ADDED
|
@@ -0,0 +1,153 @@
|
|
|
1
|
+
"""The only layer that triggers DuckDB result production."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
from typing import TYPE_CHECKING, Literal, cast
|
|
6
|
+
|
|
7
|
+
import numpy as np
|
|
8
|
+
import pandas as pd
|
|
9
|
+
import pyarrow as pa
|
|
10
|
+
|
|
11
|
+
from duckpd._typing import ParquetCompression
|
|
12
|
+
from duckpd.errors import MaterializationError
|
|
13
|
+
|
|
14
|
+
if TYPE_CHECKING:
|
|
15
|
+
from duckpd._compiler import DuckDBCompiler
|
|
16
|
+
from duckpd._logical import LogicalPlan
|
|
17
|
+
from duckpd.session import Session
|
|
18
|
+
|
|
19
|
+
|
|
20
|
+
class Executor:
|
|
21
|
+
"""Execute compiled plans and track observable execution boundaries."""
|
|
22
|
+
|
|
23
|
+
def __init__(self, session: Session, compiler: DuckDBCompiler) -> None:
|
|
24
|
+
self._session = session
|
|
25
|
+
self._compiler = compiler
|
|
26
|
+
|
|
27
|
+
def collect(self, plan: LogicalPlan) -> pd.DataFrame:
|
|
28
|
+
compiled = self._compiler.compile(plan)
|
|
29
|
+
self._session._begin_execution()
|
|
30
|
+
result = compiled.relation.to_df()
|
|
31
|
+
# Normalize DuckDB nullable integer columns to float64 if they contain nulls
|
|
32
|
+
for col in result.columns:
|
|
33
|
+
dtype_str = str(result[col].dtype)
|
|
34
|
+
if (dtype_str.startswith("Int") or dtype_str.startswith("UInt")) and result[
|
|
35
|
+
col
|
|
36
|
+
].isna().any():
|
|
37
|
+
result[col] = result[col].astype("float64")
|
|
38
|
+
|
|
39
|
+
index_ids = plan.metadata.index.columns
|
|
40
|
+
if index_ids:
|
|
41
|
+
index_labels = [compiled.bindings[column_id] for column_id in index_ids]
|
|
42
|
+
result = result.set_index(index_labels, drop=plan.metadata.index.drop)
|
|
43
|
+
hidden_labels = [
|
|
44
|
+
compiled.bindings[column.id]
|
|
45
|
+
for column in plan.metadata.columns
|
|
46
|
+
if column.hidden and column.id not in index_ids
|
|
47
|
+
]
|
|
48
|
+
if hidden_labels:
|
|
49
|
+
result = result.drop(columns=hidden_labels)
|
|
50
|
+
return result
|
|
51
|
+
|
|
52
|
+
def to_arrow(self, plan: LogicalPlan) -> pa.Table:
|
|
53
|
+
compiled = self._compiler.compile(plan)
|
|
54
|
+
self._session._begin_execution()
|
|
55
|
+
return self._compiler.project_visible(compiled, plan).relation.to_arrow_table()
|
|
56
|
+
|
|
57
|
+
def to_arrow_batches(
|
|
58
|
+
self, plan: LogicalPlan, *, batch_size: int
|
|
59
|
+
) -> pa.RecordBatchReader:
|
|
60
|
+
compiled = self._compiler.compile(plan)
|
|
61
|
+
self._session._begin_execution()
|
|
62
|
+
return self._compiler.project_visible(compiled, plan).relation.to_arrow_reader(
|
|
63
|
+
batch_size
|
|
64
|
+
)
|
|
65
|
+
|
|
66
|
+
def write_parquet(
|
|
67
|
+
self,
|
|
68
|
+
plan: LogicalPlan,
|
|
69
|
+
path: str,
|
|
70
|
+
*,
|
|
71
|
+
compression: ParquetCompression,
|
|
72
|
+
overwrite: bool,
|
|
73
|
+
) -> None:
|
|
74
|
+
compiled = self._compiler.compile(plan)
|
|
75
|
+
self._session._begin_execution()
|
|
76
|
+
self._compiler.project_visible(compiled, plan).relation.write_parquet(
|
|
77
|
+
path,
|
|
78
|
+
compression=compression,
|
|
79
|
+
overwrite=overwrite,
|
|
80
|
+
)
|
|
81
|
+
|
|
82
|
+
def explain(
|
|
83
|
+
self,
|
|
84
|
+
plan: LogicalPlan,
|
|
85
|
+
*,
|
|
86
|
+
mode: Literal["all", "logical", "sql", "physical"] = "all",
|
|
87
|
+
) -> str:
|
|
88
|
+
compiled = self._compiler.compile(plan)
|
|
89
|
+
relation = compiled.relation
|
|
90
|
+
self._session._begin_execution()
|
|
91
|
+
if mode == "logical":
|
|
92
|
+
return f"DuckPD logical plan:\n{plan!r}"
|
|
93
|
+
if mode == "sql":
|
|
94
|
+
return f"DuckDB SQL:\n{relation.sql_query()}"
|
|
95
|
+
if mode == "physical":
|
|
96
|
+
return f"DuckDB physical plan:\n{relation.explain()}"
|
|
97
|
+
if mode == "all":
|
|
98
|
+
return (
|
|
99
|
+
f"DuckPD logical plan:\n{plan!r}\n\n"
|
|
100
|
+
f"DuckDB SQL:\n{relation.sql_query()}\n\n"
|
|
101
|
+
f"DuckDB physical plan:\n{relation.explain()}"
|
|
102
|
+
)
|
|
103
|
+
msg = (
|
|
104
|
+
f"Unknown explain mode: {mode!r}; "
|
|
105
|
+
"expected 'all', 'logical', 'sql', or 'physical'"
|
|
106
|
+
)
|
|
107
|
+
raise ValueError(msg)
|
|
108
|
+
|
|
109
|
+
def explain_write(
|
|
110
|
+
self,
|
|
111
|
+
plan: LogicalPlan,
|
|
112
|
+
path: str,
|
|
113
|
+
*,
|
|
114
|
+
compression: ParquetCompression = "snappy",
|
|
115
|
+
) -> str:
|
|
116
|
+
"""Inspect write strategy and execution plan without writing rows."""
|
|
117
|
+
compiled = self._compiler.compile(plan)
|
|
118
|
+
visible_rel = self._compiler.project_visible(compiled, plan).relation
|
|
119
|
+
self._session._begin_execution()
|
|
120
|
+
return (
|
|
121
|
+
f"Write target: {path}\n"
|
|
122
|
+
f"Compression: {compression}\n"
|
|
123
|
+
f"Output columns: {list(plan.metadata.visible_columns)}\n"
|
|
124
|
+
f"DuckDB physical plan:\n{visible_rel.explain()}"
|
|
125
|
+
)
|
|
126
|
+
|
|
127
|
+
def reduce_scalar(self, plan: LogicalPlan) -> object:
|
|
128
|
+
"""Execute a one-column, one-row aggregate plan."""
|
|
129
|
+
compiled = self._compiler.compile(plan)
|
|
130
|
+
if len(plan.metadata.visible_columns) != 1:
|
|
131
|
+
raise MaterializationError("Scalar reduction requires one output column")
|
|
132
|
+
self._session._begin_execution()
|
|
133
|
+
result = compiled.relation.to_df()
|
|
134
|
+
if result.shape != (1, 1):
|
|
135
|
+
raise MaterializationError("Scalar reduction did not produce one value")
|
|
136
|
+
value = cast("object", result.iloc[0, 0])
|
|
137
|
+
return np.nan if value is None else value
|
|
138
|
+
|
|
139
|
+
def reduce_columns(self, plan: LogicalPlan) -> pd.Series:
|
|
140
|
+
"""Execute a one-row aggregate plan as a label-indexed pandas Series."""
|
|
141
|
+
compiled = self._compiler.compile(plan)
|
|
142
|
+
self._session._begin_execution()
|
|
143
|
+
result = compiled.relation.to_df()
|
|
144
|
+
if result.shape != (1, len(plan.metadata.visible_columns)):
|
|
145
|
+
raise MaterializationError("Column reduction did not produce one row")
|
|
146
|
+
reduced = result.iloc[0]
|
|
147
|
+
reduced.index = [column.label for column in plan.metadata.visible_columns]
|
|
148
|
+
reduced.name = None
|
|
149
|
+
if reduced.isna().all():
|
|
150
|
+
return pd.Series(np.nan, index=reduced.index, dtype="float64")
|
|
151
|
+
if reduced.isna().any():
|
|
152
|
+
reduced = reduced.map(lambda value: np.nan if value is None else value)
|
|
153
|
+
return reduced.infer_objects()
|