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.
@@ -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
+ )