calc-flow-python 4.0.0__cp313-abi3-win_amd64.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.
- calc_flow/__init__.py +169 -0
- calc_flow/_native.pyd +0 -0
- calc_flow/_native.pyi +366 -0
- calc_flow/array.py +1324 -0
- calc_flow/capabilities.py +775 -0
- calc_flow/config.py +219 -0
- calc_flow/errors.py +25 -0
- calc_flow/join_spec.py +156 -0
- calc_flow/pipeline.py +1123 -0
- calc_flow/py.typed +0 -0
- calc_flow/runtime.py +979 -0
- calc_flow/store.py +138 -0
- calc_flow/symbolic/__init__.py +65 -0
- calc_flow/symbolic/_generated_rolling_kernels.py +23 -0
- calc_flow/symbolic/analyzer.py +2280 -0
- calc_flow/symbolic/domains.py +77 -0
- calc_flow/symbolic/errors.py +58 -0
- calc_flow/symbolic/expr.py +662 -0
- calc_flow/symbolic/lower/__init__.py +31 -0
- calc_flow/symbolic/lower/planners.py +1270 -0
- calc_flow/symbolic/lower/program.py +840 -0
- calc_flow/symbolic/lower/segments.py +836 -0
- calc_flow/symbolic/lower/strategies.py +1472 -0
- calc_flow/symbolic/nodes.py +603 -0
- calc_flow/symbolic/ops.py +1155 -0
- calc_flow/symbolic/optimizer.py +600 -0
- calc_flow/symbolic/program.py +377 -0
- calc_flow/symbolic/types.py +110 -0
- calc_flow/symbolic/windows.py +153 -0
- calc_flow/udf.py +19 -0
- calc_flow_python-4.0.0.dist-info/METADATA +376 -0
- calc_flow_python-4.0.0.dist-info/RECORD +35 -0
- calc_flow_python-4.0.0.dist-info/WHEEL +4 -0
- calc_flow_python-4.0.0.dist-info/licenses/LICENSE +202 -0
- calc_flow_python-4.0.0.dist-info/sboms/calc-flow-python.cyclonedx.json +10081 -0
|
@@ -0,0 +1,836 @@
|
|
|
1
|
+
"""Segment constants, cache identity, and SQL rendering for the lowerer."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
from dataclasses import dataclass
|
|
6
|
+
from typing import TYPE_CHECKING, Final
|
|
7
|
+
|
|
8
|
+
from calc_flow.symbolic import errors
|
|
9
|
+
from calc_flow.symbolic.analyzer import (
|
|
10
|
+
_ROW_LOCAL_PRIMITIVES,
|
|
11
|
+
_literal_dtype,
|
|
12
|
+
_schema_fields,
|
|
13
|
+
)
|
|
14
|
+
from calc_flow.symbolic.nodes import (
|
|
15
|
+
CBool,
|
|
16
|
+
CDType,
|
|
17
|
+
CEnum,
|
|
18
|
+
CFloat,
|
|
19
|
+
CInt,
|
|
20
|
+
CMap,
|
|
21
|
+
CNull,
|
|
22
|
+
CSeq,
|
|
23
|
+
CStr,
|
|
24
|
+
CValue,
|
|
25
|
+
Node,
|
|
26
|
+
build,
|
|
27
|
+
)
|
|
28
|
+
from calc_flow.symbolic.types import Field
|
|
29
|
+
|
|
30
|
+
if TYPE_CHECKING:
|
|
31
|
+
pass
|
|
32
|
+
|
|
33
|
+
|
|
34
|
+
@dataclass(frozen=True, slots=True)
|
|
35
|
+
class _CompileCacheKey:
|
|
36
|
+
"""Deterministic, runtime-scoped identity for one symbolic compilation."""
|
|
37
|
+
|
|
38
|
+
program_fingerprint: str
|
|
39
|
+
mode: str
|
|
40
|
+
input_declarations: tuple[str, ...]
|
|
41
|
+
capability_schema_version: int
|
|
42
|
+
capability_session_id: str
|
|
43
|
+
capability_revision: int
|
|
44
|
+
operator_versions: tuple[tuple[str, str], ...]
|
|
45
|
+
provider_versions: tuple[tuple[str, str, str], ...]
|
|
46
|
+
udf_versions: tuple[tuple[str, str, str], ...]
|
|
47
|
+
allowed_lateness_micros: int
|
|
48
|
+
late_policy: str
|
|
49
|
+
|
|
50
|
+
|
|
51
|
+
_TABLE_OUTPUT_PRIMITIVES: Final = frozenset(
|
|
52
|
+
{"table_input", "project", "filter", "with_columns"}
|
|
53
|
+
)
|
|
54
|
+
|
|
55
|
+
_MATRIX_PRIMITIVES: Final = frozenset(
|
|
56
|
+
{
|
|
57
|
+
"add",
|
|
58
|
+
"and",
|
|
59
|
+
"eq",
|
|
60
|
+
"ge",
|
|
61
|
+
"gt",
|
|
62
|
+
"le",
|
|
63
|
+
"lt",
|
|
64
|
+
"matmul",
|
|
65
|
+
"mul",
|
|
66
|
+
"ne",
|
|
67
|
+
"neg",
|
|
68
|
+
"not",
|
|
69
|
+
"or",
|
|
70
|
+
"sub",
|
|
71
|
+
"truediv",
|
|
72
|
+
}
|
|
73
|
+
)
|
|
74
|
+
|
|
75
|
+
_ROLLING_PRIMITIVES: Final = frozenset(
|
|
76
|
+
{
|
|
77
|
+
"lag",
|
|
78
|
+
"delta",
|
|
79
|
+
"ewma",
|
|
80
|
+
"count",
|
|
81
|
+
"sum",
|
|
82
|
+
"mean",
|
|
83
|
+
"min",
|
|
84
|
+
"max",
|
|
85
|
+
"variance",
|
|
86
|
+
"stddev",
|
|
87
|
+
"covariance",
|
|
88
|
+
"correlation",
|
|
89
|
+
}
|
|
90
|
+
)
|
|
91
|
+
|
|
92
|
+
_ROLLING_DDOF_PRIMITIVES: Final = frozenset(
|
|
93
|
+
{"variance", "stddev", "covariance", "correlation"}
|
|
94
|
+
)
|
|
95
|
+
|
|
96
|
+
_ROLLING_PAIR_PRIMITIVES: Final = frozenset({"covariance", "correlation"})
|
|
97
|
+
|
|
98
|
+
_CROSS_SECTION_PRIMITIVES: Final = frozenset(
|
|
99
|
+
{
|
|
100
|
+
"rank",
|
|
101
|
+
"percentile",
|
|
102
|
+
"demean",
|
|
103
|
+
"zscore",
|
|
104
|
+
"winsorize",
|
|
105
|
+
"top",
|
|
106
|
+
"bottom",
|
|
107
|
+
"mean_fill",
|
|
108
|
+
}
|
|
109
|
+
)
|
|
110
|
+
|
|
111
|
+
_CROSS_SECTION_ORDERING: Final = ("rank", "percentile")
|
|
112
|
+
|
|
113
|
+
_CROSS_SECTION_DDOF: Final = "zscore"
|
|
114
|
+
|
|
115
|
+
_U64_MAX: Final = (1 << 64) - 1
|
|
116
|
+
|
|
117
|
+
_BINARY_SQL: Final = {
|
|
118
|
+
"add": "+",
|
|
119
|
+
"sub": "-",
|
|
120
|
+
"mul": "*",
|
|
121
|
+
"truediv": "/",
|
|
122
|
+
"eq": "=",
|
|
123
|
+
"ne": "!=",
|
|
124
|
+
"lt": "<",
|
|
125
|
+
"le": "<=",
|
|
126
|
+
"gt": ">",
|
|
127
|
+
"ge": ">=",
|
|
128
|
+
"and": "AND",
|
|
129
|
+
"or": "OR",
|
|
130
|
+
}
|
|
131
|
+
|
|
132
|
+
_FUNCTION_SQL: Final = {"log": "ln", "exp": "exp", "sqrt": "sqrt", "abs": "abs"}
|
|
133
|
+
|
|
134
|
+
_CAST_TYPES: Final = {
|
|
135
|
+
"bool": "BOOLEAN",
|
|
136
|
+
"int8": "TINYINT",
|
|
137
|
+
"int16": "SMALLINT",
|
|
138
|
+
"int32": "INT",
|
|
139
|
+
"int64": "BIGINT",
|
|
140
|
+
"uint8": "TINYINT UNSIGNED",
|
|
141
|
+
"uint16": "SMALLINT UNSIGNED",
|
|
142
|
+
"uint32": "INT UNSIGNED",
|
|
143
|
+
"uint64": "BIGINT UNSIGNED",
|
|
144
|
+
"float32": "REAL",
|
|
145
|
+
"float64": "DOUBLE",
|
|
146
|
+
}
|
|
147
|
+
|
|
148
|
+
|
|
149
|
+
@dataclass(frozen=True, slots=True)
|
|
150
|
+
class _Segment:
|
|
151
|
+
"""One fused row-local table resolution over one table input lineage.
|
|
152
|
+
|
|
153
|
+
``predicate`` filters declared *below* every rolling feature (they feed
|
|
154
|
+
the rolling stage); ``post_predicate`` filters declared *above* them
|
|
155
|
+
(they apply after). Without rolling primitives both fuse at the final
|
|
156
|
+
stage, preserving the historical behavior.
|
|
157
|
+
"""
|
|
158
|
+
|
|
159
|
+
input_node: Node
|
|
160
|
+
fields: tuple[str, ...]
|
|
161
|
+
env: tuple[tuple[str, Node], ...]
|
|
162
|
+
predicate: Node | None
|
|
163
|
+
post_predicate: Node | None = None
|
|
164
|
+
|
|
165
|
+
|
|
166
|
+
def _cstr(value: CValue | None, /) -> str:
|
|
167
|
+
return value.value if isinstance(value, CStr) else ""
|
|
168
|
+
|
|
169
|
+
|
|
170
|
+
def _cstr_seq(value: CValue | None, /) -> tuple[str, ...]:
|
|
171
|
+
|
|
172
|
+
if isinstance(value, CSeq):
|
|
173
|
+
return tuple(item.value for item in value.items if isinstance(item, CStr))
|
|
174
|
+
return ()
|
|
175
|
+
|
|
176
|
+
|
|
177
|
+
def _base_ref(name: str, /) -> Node:
|
|
178
|
+
return build("column_ref", (), {"name": CStr(name)})
|
|
179
|
+
|
|
180
|
+
|
|
181
|
+
def _reject_primitive(path: str, node: Node, /) -> None:
|
|
182
|
+
errors.raise_compile(
|
|
183
|
+
path,
|
|
184
|
+
errors.UNKNOWN_PRIMITIVE_VERSION,
|
|
185
|
+
f"primitive {node.op.name!r} is not supported by the row-local lowerer",
|
|
186
|
+
)
|
|
187
|
+
|
|
188
|
+
|
|
189
|
+
def _resolve_table(node: Node, path: str, /) -> _Segment:
|
|
190
|
+
name = node.op.name
|
|
191
|
+
if name == "table_input":
|
|
192
|
+
return _resolve_table_input(node)
|
|
193
|
+
if name == "project":
|
|
194
|
+
return _resolve_project(node, path)
|
|
195
|
+
if name == "filter":
|
|
196
|
+
return _resolve_filter(node, path)
|
|
197
|
+
if name == "with_columns":
|
|
198
|
+
return _resolve_with_columns(node, path)
|
|
199
|
+
_reject_primitive(path, node)
|
|
200
|
+
|
|
201
|
+
|
|
202
|
+
def _resolve_table_input(node: Node, /) -> _Segment:
|
|
203
|
+
fields = _schema_fields(node.attr("schema"))
|
|
204
|
+
return _Segment(
|
|
205
|
+
node,
|
|
206
|
+
tuple(field.name for field in fields),
|
|
207
|
+
tuple((field.name, _base_ref(field.name)) for field in fields),
|
|
208
|
+
None,
|
|
209
|
+
)
|
|
210
|
+
|
|
211
|
+
|
|
212
|
+
def _resolve_project(node: Node, path: str, /) -> _Segment:
|
|
213
|
+
child = _resolve_table(node.args[0], f"{path}.project.value")
|
|
214
|
+
env = dict(child.env)
|
|
215
|
+
columns = _cstr_seq(node.attr("columns"))
|
|
216
|
+
return _Segment(
|
|
217
|
+
child.input_node,
|
|
218
|
+
columns,
|
|
219
|
+
tuple((field, env[field]) for field in columns),
|
|
220
|
+
child.predicate,
|
|
221
|
+
child.post_predicate,
|
|
222
|
+
)
|
|
223
|
+
|
|
224
|
+
|
|
225
|
+
def _filter_needs_stateful_tail(child: _Segment, predicate: Node, /) -> bool:
|
|
226
|
+
return (
|
|
227
|
+
_segment_has_rolling(child)
|
|
228
|
+
or any(True for _ in _find_rolling(predicate))
|
|
229
|
+
or _segment_has_cross_section(child)
|
|
230
|
+
or any(True for _ in _find_cross_section(predicate))
|
|
231
|
+
)
|
|
232
|
+
|
|
233
|
+
|
|
234
|
+
def _resolve_filter(node: Node, path: str, /) -> _Segment:
|
|
235
|
+
child = _resolve_table(node.args[0], f"{path}.filter.value")
|
|
236
|
+
predicate = _inline(node.args[1], dict(child.env), f"{path}.filter.predicate")
|
|
237
|
+
if _filter_needs_stateful_tail(child, predicate):
|
|
238
|
+
combined = (
|
|
239
|
+
predicate
|
|
240
|
+
if child.post_predicate is None
|
|
241
|
+
else build("and", (child.post_predicate, predicate), {})
|
|
242
|
+
)
|
|
243
|
+
return _Segment(
|
|
244
|
+
child.input_node, child.fields, child.env, child.predicate, combined
|
|
245
|
+
)
|
|
246
|
+
combined = (
|
|
247
|
+
predicate
|
|
248
|
+
if child.predicate is None
|
|
249
|
+
else build("and", (child.predicate, predicate), {})
|
|
250
|
+
)
|
|
251
|
+
return _Segment(
|
|
252
|
+
child.input_node, child.fields, child.env, combined, child.post_predicate
|
|
253
|
+
)
|
|
254
|
+
|
|
255
|
+
|
|
256
|
+
def _resolve_with_columns(node: Node, path: str, /) -> _Segment:
|
|
257
|
+
child = _resolve_table(node.args[0], f"{path}.with_columns.value")
|
|
258
|
+
env = dict(child.env)
|
|
259
|
+
names = _cstr_seq(node.attr("names"))
|
|
260
|
+
for index, feature in enumerate(names):
|
|
261
|
+
env[feature] = _inline(node.args[index + 1], env, f"{path}.{feature}")
|
|
262
|
+
return _Segment(
|
|
263
|
+
child.input_node,
|
|
264
|
+
(*child.fields, *names),
|
|
265
|
+
tuple(env.items()),
|
|
266
|
+
child.predicate,
|
|
267
|
+
child.post_predicate,
|
|
268
|
+
)
|
|
269
|
+
|
|
270
|
+
|
|
271
|
+
def _segment_has_rolling(segment: _Segment, /) -> bool:
|
|
272
|
+
return any(True for _, tree in segment.env for _ in _find_rolling(tree)) or any(
|
|
273
|
+
True
|
|
274
|
+
for tree in (segment.predicate, segment.post_predicate)
|
|
275
|
+
if tree is not None
|
|
276
|
+
for _ in _find_rolling(tree)
|
|
277
|
+
)
|
|
278
|
+
|
|
279
|
+
|
|
280
|
+
def _segment_has_cross_section(segment: _Segment, /) -> bool:
|
|
281
|
+
return any(
|
|
282
|
+
True for _, tree in segment.env for _ in _find_cross_section(tree)
|
|
283
|
+
) or any(
|
|
284
|
+
True
|
|
285
|
+
for tree in (segment.predicate, segment.post_predicate)
|
|
286
|
+
if tree is not None
|
|
287
|
+
for _ in _find_cross_section(tree)
|
|
288
|
+
)
|
|
289
|
+
|
|
290
|
+
|
|
291
|
+
def _inline(node: Node, env: dict[str, Node], path: str, /) -> Node:
|
|
292
|
+
name = node.op.name
|
|
293
|
+
if name == "column_ref":
|
|
294
|
+
return env[_cstr(node.attr("name"))]
|
|
295
|
+
if name == "literal":
|
|
296
|
+
return node
|
|
297
|
+
if (
|
|
298
|
+
name not in _ROW_LOCAL_PRIMITIVES
|
|
299
|
+
and name not in _ROLLING_PRIMITIVES
|
|
300
|
+
and name not in _CROSS_SECTION_PRIMITIVES
|
|
301
|
+
):
|
|
302
|
+
_reject_primitive(path, node)
|
|
303
|
+
if name == "cast":
|
|
304
|
+
_cast_target(node, path)
|
|
305
|
+
return build(
|
|
306
|
+
name,
|
|
307
|
+
tuple(_inline(argument, env, path) for argument in node.args),
|
|
308
|
+
dict(node.attrs.entries),
|
|
309
|
+
version=node.op.version,
|
|
310
|
+
)
|
|
311
|
+
|
|
312
|
+
|
|
313
|
+
def _find_primitives(node: Node, primitives: frozenset[str], /):
|
|
314
|
+
"""Yield matching subtrees in deterministic first-appearance order."""
|
|
315
|
+
|
|
316
|
+
if node.op.name in primitives:
|
|
317
|
+
yield node
|
|
318
|
+
for argument in node.args:
|
|
319
|
+
yield from _find_primitives(argument, primitives)
|
|
320
|
+
|
|
321
|
+
|
|
322
|
+
def _find_rolling(node: Node, /):
|
|
323
|
+
"""Yield every rolling temporal subtree in first-appearance order."""
|
|
324
|
+
|
|
325
|
+
yield from _find_primitives(node, _ROLLING_PRIMITIVES)
|
|
326
|
+
|
|
327
|
+
|
|
328
|
+
def _find_cross_section(node: Node, /):
|
|
329
|
+
"""Yield every cross-section subtree in first-appearance order."""
|
|
330
|
+
|
|
331
|
+
yield from _find_primitives(node, _CROSS_SECTION_PRIMITIVES)
|
|
332
|
+
|
|
333
|
+
|
|
334
|
+
def _cast_target(node: Node, path: str, /) -> str:
|
|
335
|
+
raw = node.attr("data_type")
|
|
336
|
+
declared = _cstr(raw) or (raw.name if isinstance(raw, CDType) else "")
|
|
337
|
+
target = _CAST_TYPES.get(declared)
|
|
338
|
+
if target is None:
|
|
339
|
+
errors.raise_compile(
|
|
340
|
+
f"{path}.cast.data_type",
|
|
341
|
+
errors.UNSUPPORTED_TYPE,
|
|
342
|
+
f"cast target {declared!r} is not portable in the row-local lowerer",
|
|
343
|
+
)
|
|
344
|
+
return target
|
|
345
|
+
|
|
346
|
+
|
|
347
|
+
def _sql_operator(name: str, node: Node, /) -> str | None:
|
|
348
|
+
"""Render the single-operand and fixed-shape SQL operators."""
|
|
349
|
+
if name == "column_ref":
|
|
350
|
+
return _quote_identifier(_cstr(node.attr("name")))
|
|
351
|
+
if name == "literal":
|
|
352
|
+
return _sql_literal(node.attr("value"))
|
|
353
|
+
if name in _BINARY_SQL:
|
|
354
|
+
return f"({_sql(node.args[0])} {_BINARY_SQL[name]} {_sql(node.args[1])})"
|
|
355
|
+
if name == "neg":
|
|
356
|
+
return f"(-{_sql(node.args[0])})"
|
|
357
|
+
if name == "not":
|
|
358
|
+
return f"(NOT {_sql(node.args[0])})"
|
|
359
|
+
return None
|
|
360
|
+
|
|
361
|
+
|
|
362
|
+
def _sql(node: Node, /) -> str:
|
|
363
|
+
name = node.op.name
|
|
364
|
+
simple = _sql_operator(name, node)
|
|
365
|
+
if simple is not None:
|
|
366
|
+
return simple
|
|
367
|
+
if name == "where":
|
|
368
|
+
return (
|
|
369
|
+
f"(CASE WHEN {_sql(node.args[0])} THEN {_sql(node.args[1])}"
|
|
370
|
+
f" ELSE {_sql(node.args[2])} END)"
|
|
371
|
+
)
|
|
372
|
+
if name == "coalesce":
|
|
373
|
+
return "COALESCE(" + ", ".join(_sql(argument) for argument in node.args) + ")"
|
|
374
|
+
if name in _FUNCTION_SQL:
|
|
375
|
+
return f"{_FUNCTION_SQL[name]}({_sql(node.args[0])})"
|
|
376
|
+
if name == "clip":
|
|
377
|
+
value = _sql(node.args[0])
|
|
378
|
+
lower = _sql_literal(node.attr("lower"))
|
|
379
|
+
upper = _sql_literal(node.attr("upper"))
|
|
380
|
+
return (
|
|
381
|
+
f"(CASE WHEN {value} < {lower} THEN {lower}"
|
|
382
|
+
f" WHEN {value} > {upper} THEN {upper} ELSE {value} END)"
|
|
383
|
+
)
|
|
384
|
+
if name == "cast":
|
|
385
|
+
return f"CAST({_sql(node.args[0])} AS {_CAST_TYPES[_cast_type_name(node)]})"
|
|
386
|
+
raise AssertionError(f"unlowerable primitive reached SQL rendering: {name}")
|
|
387
|
+
|
|
388
|
+
|
|
389
|
+
def _cast_type_name(node: Node, /) -> str:
|
|
390
|
+
raw = node.attr("data_type")
|
|
391
|
+
return _cstr(raw) or (raw.name if isinstance(raw, CDType) else "")
|
|
392
|
+
|
|
393
|
+
|
|
394
|
+
def _sql_literal(value: CValue | None, /) -> str:
|
|
395
|
+
if isinstance(value, CNull) or value is None:
|
|
396
|
+
return "NULL"
|
|
397
|
+
if isinstance(value, CBool):
|
|
398
|
+
return "TRUE" if value.value else "FALSE"
|
|
399
|
+
if isinstance(value, CInt):
|
|
400
|
+
return str(value.value)
|
|
401
|
+
if isinstance(value, CFloat):
|
|
402
|
+
return repr(value.value)
|
|
403
|
+
if isinstance(value, CStr):
|
|
404
|
+
return "'" + value.value.replace("'", "''") + "'"
|
|
405
|
+
raise AssertionError(f"unsupported literal reached SQL rendering: {value!r}")
|
|
406
|
+
|
|
407
|
+
|
|
408
|
+
def _quote_identifier(name: str, /) -> str:
|
|
409
|
+
return '"' + name.replace('"', '""') + '"'
|
|
410
|
+
|
|
411
|
+
|
|
412
|
+
def _select_item(name: str, tree: Node, /) -> str:
|
|
413
|
+
if tree.op.name == "column_ref" and tree.attr("name") == CStr(name):
|
|
414
|
+
return _quote_identifier(name)
|
|
415
|
+
return f"{_sql(tree)} AS {_quote_identifier(name)}"
|
|
416
|
+
|
|
417
|
+
|
|
418
|
+
def _expression_node(
|
|
419
|
+
node_id: str,
|
|
420
|
+
select: list[str],
|
|
421
|
+
filter_sql: str | None,
|
|
422
|
+
input_schema: tuple[Field, ...] | None,
|
|
423
|
+
output_schema: tuple[Field, ...] | None = None,
|
|
424
|
+
/,
|
|
425
|
+
) -> dict[str, object]:
|
|
426
|
+
node: dict[str, object] = {
|
|
427
|
+
"id": node_id,
|
|
428
|
+
"operator": {
|
|
429
|
+
"kind": "expression",
|
|
430
|
+
"expression": "",
|
|
431
|
+
"select": select,
|
|
432
|
+
"filter": filter_sql,
|
|
433
|
+
"udfs": [],
|
|
434
|
+
},
|
|
435
|
+
}
|
|
436
|
+
if input_schema is not None:
|
|
437
|
+
node["input_ports"] = [
|
|
438
|
+
{
|
|
439
|
+
"name": "input",
|
|
440
|
+
"kind": "table",
|
|
441
|
+
"required": True,
|
|
442
|
+
"schema": [_field_json(field) for field in input_schema],
|
|
443
|
+
}
|
|
444
|
+
]
|
|
445
|
+
if output_schema is not None:
|
|
446
|
+
node["output_ports"] = [
|
|
447
|
+
{
|
|
448
|
+
"name": "output",
|
|
449
|
+
"kind": "table",
|
|
450
|
+
"required": True,
|
|
451
|
+
"schema": [_field_json(field) for field in output_schema],
|
|
452
|
+
}
|
|
453
|
+
]
|
|
454
|
+
return node
|
|
455
|
+
|
|
456
|
+
|
|
457
|
+
def _field_json(field: Field, /) -> dict[str, object]:
|
|
458
|
+
return {
|
|
459
|
+
"name": field.name,
|
|
460
|
+
"data_type": field.data_type,
|
|
461
|
+
"nullable": field.nullable,
|
|
462
|
+
}
|
|
463
|
+
|
|
464
|
+
|
|
465
|
+
def _cint(value: CValue | None, /) -> int | None:
|
|
466
|
+
return value.value if isinstance(value, CInt) else None
|
|
467
|
+
|
|
468
|
+
|
|
469
|
+
def _cnumber(value: CValue | None, /) -> int | float | None:
|
|
470
|
+
return value.value if isinstance(value, (CInt, CFloat)) else None
|
|
471
|
+
|
|
472
|
+
|
|
473
|
+
def _cbool(value: CValue | None, /) -> bool | None:
|
|
474
|
+
return value.value if isinstance(value, CBool) else None
|
|
475
|
+
|
|
476
|
+
|
|
477
|
+
def _replace_materialized(node: Node, replacements: dict[str, str], /) -> Node:
|
|
478
|
+
replacement = replacements.get(node.digest)
|
|
479
|
+
if replacement is not None:
|
|
480
|
+
return _base_ref(replacement)
|
|
481
|
+
return build(
|
|
482
|
+
node.op.name,
|
|
483
|
+
tuple(_replace_materialized(argument, replacements) for argument in node.args),
|
|
484
|
+
dict(node.attrs.entries),
|
|
485
|
+
version=node.op.version,
|
|
486
|
+
)
|
|
487
|
+
|
|
488
|
+
|
|
489
|
+
@dataclass(frozen=True, slots=True)
|
|
490
|
+
class _RollingPlan:
|
|
491
|
+
"""One lowered rolling stage: the project node plus the rewritten
|
|
492
|
+
row-local environment that references its output columns."""
|
|
493
|
+
|
|
494
|
+
node_id: str
|
|
495
|
+
node: dict[str, object]
|
|
496
|
+
materialization_node_id: str | None
|
|
497
|
+
materialization_node: dict[str, object] | None
|
|
498
|
+
env: tuple[tuple[str, Node], ...]
|
|
499
|
+
post_predicate: Node | None
|
|
500
|
+
input_field_names: tuple[str, ...]
|
|
501
|
+
output_fields: tuple[Field, ...]
|
|
502
|
+
replacements: tuple[tuple[str, str], ...]
|
|
503
|
+
|
|
504
|
+
|
|
505
|
+
@dataclass(frozen=True, slots=True)
|
|
506
|
+
class _RollingPipeline:
|
|
507
|
+
"""Ordered rolling stages and their final rewritten environment."""
|
|
508
|
+
|
|
509
|
+
stages: tuple[_RollingPlan, ...]
|
|
510
|
+
env: tuple[tuple[str, Node], ...]
|
|
511
|
+
post_predicate: Node | None
|
|
512
|
+
input_field_names: tuple[str, ...]
|
|
513
|
+
output_fields: tuple[Field, ...]
|
|
514
|
+
|
|
515
|
+
@property
|
|
516
|
+
def node_id(self) -> str:
|
|
517
|
+
"""Return the final state stage identifier."""
|
|
518
|
+
|
|
519
|
+
return self.stages[-1].node_id
|
|
520
|
+
|
|
521
|
+
|
|
522
|
+
@dataclass(frozen=True, slots=True)
|
|
523
|
+
class _StatefulInputPlan:
|
|
524
|
+
"""One deterministic row-local materialization before native state."""
|
|
525
|
+
|
|
526
|
+
names: tuple[tuple[str, str], ...]
|
|
527
|
+
fields: tuple[Field, ...]
|
|
528
|
+
node_id: str | None
|
|
529
|
+
node: dict[str, object] | None
|
|
530
|
+
input_fields: tuple[Field, ...]
|
|
531
|
+
used_names: frozenset[str]
|
|
532
|
+
|
|
533
|
+
|
|
534
|
+
@dataclass(frozen=True, slots=True)
|
|
535
|
+
class _StatefulInputRequest:
|
|
536
|
+
output_name: str
|
|
537
|
+
path: str
|
|
538
|
+
column_stem: str
|
|
539
|
+
node_id: str
|
|
540
|
+
domain: str
|
|
541
|
+
|
|
542
|
+
|
|
543
|
+
def _required_stateful_input(
|
|
544
|
+
request: _StatefulInputRequest,
|
|
545
|
+
primitive: str,
|
|
546
|
+
argument: Node,
|
|
547
|
+
index: int,
|
|
548
|
+
input_types: dict[str, Field],
|
|
549
|
+
reserved: set[str],
|
|
550
|
+
/,
|
|
551
|
+
) -> tuple[str, Field, str]:
|
|
552
|
+
if not _rolling_argument_is_row_local(argument):
|
|
553
|
+
errors.raise_compile(
|
|
554
|
+
request.path,
|
|
555
|
+
errors.UNSUPPORTED_TYPE,
|
|
556
|
+
f"{request.domain} {primitive} argument must be an input column"
|
|
557
|
+
" or row-local expression after earlier state staging",
|
|
558
|
+
)
|
|
559
|
+
name = f"{request.output_name}__cf_{request.column_stem}_{index}"
|
|
560
|
+
if name in reserved:
|
|
561
|
+
errors.raise_compile(
|
|
562
|
+
f"{request.path}.{name}",
|
|
563
|
+
errors.DUPLICATE_NAME,
|
|
564
|
+
f"materialized {request.domain} input {name!r} collides"
|
|
565
|
+
" with a declared field",
|
|
566
|
+
)
|
|
567
|
+
return (
|
|
568
|
+
name,
|
|
569
|
+
_row_local_field(argument, name, input_types),
|
|
570
|
+
_select_item(name, argument),
|
|
571
|
+
)
|
|
572
|
+
|
|
573
|
+
|
|
574
|
+
def _plan_stateful_inputs(
|
|
575
|
+
request: _StatefulInputRequest,
|
|
576
|
+
input_fields: tuple[Field, ...],
|
|
577
|
+
used_names: set[str],
|
|
578
|
+
arguments: tuple[tuple[str, Node], ...],
|
|
579
|
+
/,
|
|
580
|
+
) -> _StatefulInputPlan:
|
|
581
|
+
input_types = {field.name: field for field in input_fields}
|
|
582
|
+
reserved = set(used_names)
|
|
583
|
+
names: dict[str, str] = {}
|
|
584
|
+
fields: list[Field] = []
|
|
585
|
+
selects = [_quote_identifier(field.name) for field in input_fields]
|
|
586
|
+
for primitive, argument in arguments:
|
|
587
|
+
if argument.op.name == "column_ref" or argument.digest in names:
|
|
588
|
+
continue
|
|
589
|
+
name, field, select = _required_stateful_input(
|
|
590
|
+
request,
|
|
591
|
+
primitive,
|
|
592
|
+
argument,
|
|
593
|
+
len(names),
|
|
594
|
+
input_types,
|
|
595
|
+
reserved,
|
|
596
|
+
)
|
|
597
|
+
reserved.add(name)
|
|
598
|
+
names[argument.digest] = name
|
|
599
|
+
fields.append(field)
|
|
600
|
+
selects.append(select)
|
|
601
|
+
state_input_fields = (*input_fields, *fields)
|
|
602
|
+
materialization_id = request.node_id if fields else None
|
|
603
|
+
materialization = (
|
|
604
|
+
_expression_node(
|
|
605
|
+
request.node_id,
|
|
606
|
+
selects,
|
|
607
|
+
None,
|
|
608
|
+
input_fields,
|
|
609
|
+
state_input_fields,
|
|
610
|
+
)
|
|
611
|
+
if fields
|
|
612
|
+
else None
|
|
613
|
+
)
|
|
614
|
+
return _StatefulInputPlan(
|
|
615
|
+
tuple(names.items()),
|
|
616
|
+
tuple(fields),
|
|
617
|
+
materialization_id,
|
|
618
|
+
materialization,
|
|
619
|
+
state_input_fields,
|
|
620
|
+
frozenset(reserved),
|
|
621
|
+
)
|
|
622
|
+
|
|
623
|
+
|
|
624
|
+
def _row_local_field(
|
|
625
|
+
node: Node,
|
|
626
|
+
name: str,
|
|
627
|
+
input_types: dict[str, Field],
|
|
628
|
+
/,
|
|
629
|
+
) -> Field:
|
|
630
|
+
"""Infer a validated row-local expression field for stateful staging."""
|
|
631
|
+
|
|
632
|
+
leaf = _row_local_leaf_field(node, name, input_types)
|
|
633
|
+
if leaf is not None:
|
|
634
|
+
return leaf
|
|
635
|
+
children = [_row_local_field(argument, name, input_types) for argument in node.args]
|
|
636
|
+
return _row_local_composite_field(node, name, children)
|
|
637
|
+
|
|
638
|
+
|
|
639
|
+
def _row_local_leaf_field(
|
|
640
|
+
node: Node,
|
|
641
|
+
name: str,
|
|
642
|
+
input_types: dict[str, Field],
|
|
643
|
+
/,
|
|
644
|
+
) -> Field | None:
|
|
645
|
+
operation = node.op.name
|
|
646
|
+
if operation == "column_ref":
|
|
647
|
+
source = input_types[_cstr(node.attr("name"))]
|
|
648
|
+
return Field(name, source.data_type, nullable=source.nullable)
|
|
649
|
+
if operation != "literal":
|
|
650
|
+
return None
|
|
651
|
+
value = node.attr("value")
|
|
652
|
+
data_type = None if value is None else _literal_dtype(value)
|
|
653
|
+
if data_type is None:
|
|
654
|
+
raise RuntimeError("validated rolling literal has no data type")
|
|
655
|
+
return Field(name, data_type, nullable=isinstance(value, CNull))
|
|
656
|
+
|
|
657
|
+
|
|
658
|
+
def _row_local_composite_field(
|
|
659
|
+
node: Node,
|
|
660
|
+
name: str,
|
|
661
|
+
children: list[Field],
|
|
662
|
+
/,
|
|
663
|
+
) -> Field:
|
|
664
|
+
operation = node.op.name
|
|
665
|
+
nullable = _any_nullable(children)
|
|
666
|
+
if operation in {"eq", "ne", "lt", "le", "gt", "ge", "and", "or", "not"}:
|
|
667
|
+
return Field(name, "bool", nullable=nullable)
|
|
668
|
+
if operation in _FUNCTION_SQL:
|
|
669
|
+
return Field(name, "float64", nullable=True)
|
|
670
|
+
if operation == "cast":
|
|
671
|
+
return Field(name, _cast_type_name(node), nullable=nullable)
|
|
672
|
+
return _row_local_conditional_field(node, name, children, nullable)
|
|
673
|
+
|
|
674
|
+
|
|
675
|
+
def _any_nullable(fields: list[Field], /) -> bool:
|
|
676
|
+
return any(field.nullable for field in fields)
|
|
677
|
+
|
|
678
|
+
|
|
679
|
+
def _row_local_conditional_field(
|
|
680
|
+
node: Node,
|
|
681
|
+
name: str,
|
|
682
|
+
children: list[Field],
|
|
683
|
+
nullable: bool,
|
|
684
|
+
/,
|
|
685
|
+
) -> Field:
|
|
686
|
+
operation = node.op.name
|
|
687
|
+
if operation == "where":
|
|
688
|
+
return Field(
|
|
689
|
+
name,
|
|
690
|
+
children[1].data_type,
|
|
691
|
+
nullable=_where_result_is_nullable(node, children),
|
|
692
|
+
)
|
|
693
|
+
if operation == "coalesce":
|
|
694
|
+
return Field(
|
|
695
|
+
name,
|
|
696
|
+
children[0].data_type,
|
|
697
|
+
nullable=all(field.nullable for field in children),
|
|
698
|
+
)
|
|
699
|
+
return Field(name, children[0].data_type, nullable=nullable)
|
|
700
|
+
|
|
701
|
+
|
|
702
|
+
def _where_result_is_nullable(node: Node, children: list[Field], /) -> bool:
|
|
703
|
+
"""Mirror DataFusion's non-null proof for a directly guarded column."""
|
|
704
|
+
|
|
705
|
+
if children[2].nullable:
|
|
706
|
+
return True
|
|
707
|
+
selected = node.args[1]
|
|
708
|
+
condition = node.args[0]
|
|
709
|
+
if selected.op.name == "column_ref" and any(
|
|
710
|
+
argument.digest == selected.digest for argument in condition.args
|
|
711
|
+
):
|
|
712
|
+
return False
|
|
713
|
+
return children[1].nullable
|
|
714
|
+
|
|
715
|
+
|
|
716
|
+
def _rolling_argument_is_row_local(node: Node, /) -> bool:
|
|
717
|
+
return node.op.name in _ROW_LOCAL_PRIMITIVES and all(
|
|
718
|
+
_rolling_argument_is_row_local(argument) for argument in node.args
|
|
719
|
+
)
|
|
720
|
+
|
|
721
|
+
|
|
722
|
+
def _find_ready_rolling(node: Node, /):
|
|
723
|
+
"""Yield innermost rolling subtrees ready for one physical stage."""
|
|
724
|
+
|
|
725
|
+
if node.op.name in _ROLLING_PRIMITIVES:
|
|
726
|
+
nested = any(True for argument in node.args for _ in _find_rolling(argument))
|
|
727
|
+
if not nested:
|
|
728
|
+
yield node
|
|
729
|
+
return
|
|
730
|
+
for argument in node.args:
|
|
731
|
+
yield from _find_ready_rolling(argument)
|
|
732
|
+
|
|
733
|
+
|
|
734
|
+
def _rolling_frame(subtree: Node, path: str, kind: str, /) -> dict[str, object]:
|
|
735
|
+
"""Render the frozen frame JSON: row-count or duration (SCE-08)."""
|
|
736
|
+
|
|
737
|
+
frame = subtree.attr("frame")
|
|
738
|
+
variant = None
|
|
739
|
+
if isinstance(frame, CMap):
|
|
740
|
+
tag = frame.get("frame")
|
|
741
|
+
if isinstance(tag, CEnum):
|
|
742
|
+
variant = tag.variant
|
|
743
|
+
if variant == "duration":
|
|
744
|
+
micros = _cint(frame.get("micros")) if isinstance(frame, CMap) else None
|
|
745
|
+
return {"kind": "duration", "micros": 1 if micros is None else micros}
|
|
746
|
+
if variant != "rows":
|
|
747
|
+
errors.raise_compile(
|
|
748
|
+
path,
|
|
749
|
+
errors.UNSUPPORTED_TYPE,
|
|
750
|
+
f"rolling {kind} requires a rows or duration frame",
|
|
751
|
+
)
|
|
752
|
+
size = _cint(frame.get("size"))
|
|
753
|
+
return {"kind": "rows", "size": 1 if size is None else size}
|
|
754
|
+
|
|
755
|
+
|
|
756
|
+
_FUSED_FLOAT_ROLLING_LEAVES: Final = frozenset({"mean", "variance", "stddev", "ewma"})
|
|
757
|
+
|
|
758
|
+
|
|
759
|
+
def _fused_difference_outputs(
|
|
760
|
+
segment: _Segment, occurrences: tuple[Node, ...], /
|
|
761
|
+
) -> tuple[tuple[str, Node], ...]:
|
|
762
|
+
"""Return final ``left - right`` expressions safe for one state stage."""
|
|
763
|
+
|
|
764
|
+
ready = {node.digest for node in occurrences}
|
|
765
|
+
return tuple(
|
|
766
|
+
(name, tree)
|
|
767
|
+
for name, tree in segment.env
|
|
768
|
+
if tree.op.name == "sub"
|
|
769
|
+
and len(tree.args) == 2
|
|
770
|
+
and all(
|
|
771
|
+
argument.op.name in _FUSED_FLOAT_ROLLING_LEAVES and argument.digest in ready
|
|
772
|
+
for argument in tree.args
|
|
773
|
+
)
|
|
774
|
+
)
|
|
775
|
+
|
|
776
|
+
|
|
777
|
+
def _rolling_input_name(
|
|
778
|
+
argument: Node, materializations: dict[str, str], path: str, /
|
|
779
|
+
) -> str:
|
|
780
|
+
if argument.op.name == "column_ref":
|
|
781
|
+
return _cstr(argument.attr("name"))
|
|
782
|
+
name = materializations.get(argument.digest)
|
|
783
|
+
if name is None:
|
|
784
|
+
errors.raise_compile(
|
|
785
|
+
path,
|
|
786
|
+
errors.SCHEMA_MISMATCH,
|
|
787
|
+
"fused rolling input was not materialized for the state stage",
|
|
788
|
+
)
|
|
789
|
+
return name
|
|
790
|
+
|
|
791
|
+
|
|
792
|
+
def _fused_float_leaf(
|
|
793
|
+
subtree: Node,
|
|
794
|
+
materializations: dict[str, str],
|
|
795
|
+
input_types: dict[str, Field],
|
|
796
|
+
path: str,
|
|
797
|
+
/,
|
|
798
|
+
) -> dict[str, object]:
|
|
799
|
+
kind = subtree.op.name
|
|
800
|
+
input_name = _rolling_input_name(subtree.args[0], materializations, path)
|
|
801
|
+
if input_name not in input_types:
|
|
802
|
+
errors.raise_compile(
|
|
803
|
+
path,
|
|
804
|
+
errors.SCHEMA_MISMATCH,
|
|
805
|
+
f"rolling {kind} argument column {input_name!r} is not in the input schema",
|
|
806
|
+
)
|
|
807
|
+
if kind == "ewma":
|
|
808
|
+
return {
|
|
809
|
+
"kind": kind,
|
|
810
|
+
"primitive_version": 1,
|
|
811
|
+
"input": input_name,
|
|
812
|
+
"span": _cint(subtree.attr("span")),
|
|
813
|
+
"min_periods": _cint(subtree.attr("min_periods")) or 1,
|
|
814
|
+
}
|
|
815
|
+
declaration: dict[str, object] = {
|
|
816
|
+
"kind": kind,
|
|
817
|
+
"primitive_version": 1,
|
|
818
|
+
"input": input_name,
|
|
819
|
+
"frame": _rolling_frame(subtree, path, kind),
|
|
820
|
+
"min_periods": _cint(subtree.attr("min_periods")) or 1,
|
|
821
|
+
}
|
|
822
|
+
if kind in _ROLLING_DDOF_PRIMITIVES:
|
|
823
|
+
ddof = _cint(subtree.attr("ddof"))
|
|
824
|
+
declaration["ddof"] = 1 if ddof is None else ddof
|
|
825
|
+
return declaration
|
|
826
|
+
|
|
827
|
+
|
|
828
|
+
def _rolling_declaration_requires_ewma(declaration: dict[str, object], /) -> bool:
|
|
829
|
+
if declaration["kind"] == "ewma":
|
|
830
|
+
return True
|
|
831
|
+
if declaration["kind"] != "difference":
|
|
832
|
+
return False
|
|
833
|
+
return any(
|
|
834
|
+
isinstance(leaf, dict) and leaf.get("kind") == "ewma"
|
|
835
|
+
for leaf in (declaration["left"], declaration["right"])
|
|
836
|
+
)
|