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,1472 @@
1
+ """Matrix, stream-join, relational-DAG, and dedup lowering strategies.
2
+
3
+ Moved verbatim from ``symbolic/lower.py``."""
4
+
5
+ from __future__ import annotations
6
+
7
+ from dataclasses import dataclass
8
+ from typing import TYPE_CHECKING, Never
9
+
10
+ from calc_flow.join_spec import (
11
+ JoinSideWire,
12
+ bounds_wire,
13
+ join_wire_spec,
14
+ limits_wire,
15
+ )
16
+ from calc_flow.pipeline import (
17
+ _canonical,
18
+ _data_sources,
19
+ )
20
+ from calc_flow.symbolic import errors
21
+ from calc_flow.symbolic.analyzer import (
22
+ TableFacts,
23
+ _Analyzer,
24
+ _schema_fields,
25
+ )
26
+ from calc_flow.symbolic.lower.planners import (
27
+ _CrossSectionPlan,
28
+ _LoweringProgram,
29
+ _LoweringValue,
30
+ _MatrixExpression,
31
+ )
32
+ from calc_flow.symbolic.lower.segments import (
33
+ _MATRIX_PRIMITIVES,
34
+ _cint,
35
+ _cstr,
36
+ _cstr_seq,
37
+ _expression_node,
38
+ _field_json,
39
+ _quote_identifier,
40
+ _reject_primitive,
41
+ _RollingPipeline,
42
+ )
43
+ from calc_flow.symbolic.nodes import (
44
+ CBool,
45
+ CDType,
46
+ CFloat,
47
+ CInt,
48
+ CSeq,
49
+ Node,
50
+ build,
51
+ )
52
+ from calc_flow.symbolic.types import Field
53
+
54
+ if TYPE_CHECKING:
55
+ from calc_flow.symbolic.program import Program
56
+
57
+
58
+ def _matrix_literal(node: Node, path: str, /) -> bool | int | float:
59
+ value = node.attr("value")
60
+ if isinstance(value, (CBool, CInt, CFloat)):
61
+ return value.value
62
+ errors.raise_compile(
63
+ path,
64
+ errors.UNSUPPORTED_TYPE,
65
+ "symbolic matrix literals must be finite bool, int, or float values",
66
+ )
67
+
68
+
69
+ def _matrix_backend(
70
+ left: _MatrixExpression,
71
+ right: _MatrixExpression,
72
+ path: str,
73
+ /,
74
+ ) -> str:
75
+ if left.backend and right.backend and left.backend != right.backend:
76
+ errors.raise_compile(
77
+ path,
78
+ errors.CAPABILITY_MISMATCH,
79
+ "symbolic matrix operands must use one provider backend",
80
+ )
81
+ return left.backend or right.backend
82
+
83
+
84
+ def _matrix_columns(
85
+ left: _MatrixExpression,
86
+ right: _MatrixExpression,
87
+ path: str,
88
+ /,
89
+ ) -> tuple[str, ...]:
90
+ if left.columns and right.columns and left.columns != right.columns:
91
+ errors.raise_compile(
92
+ path,
93
+ errors.SCHEMA_MISMATCH,
94
+ "symbolic matrix operands must use one ordered column selection",
95
+ )
96
+ return left.columns or right.columns
97
+
98
+
99
+ def _matrix_parameter(
100
+ left: _MatrixExpression,
101
+ right: _MatrixExpression,
102
+ path: str,
103
+ /,
104
+ ) -> Node | None:
105
+ if (
106
+ left.parameter is not None
107
+ and right.parameter is not None
108
+ and left.parameter.digest != right.parameter.digest
109
+ ):
110
+ errors.raise_compile(
111
+ path,
112
+ errors.CAPABILITY_MISMATCH,
113
+ "one symbolic matrix output supports exactly one static parameter",
114
+ )
115
+ return left.parameter or right.parameter
116
+
117
+
118
+ def _merge_matrix_expression(
119
+ left: _MatrixExpression,
120
+ right: _MatrixExpression,
121
+ path: str,
122
+ /,
123
+ ) -> tuple[str, tuple[str, ...], Node | None]:
124
+ return (
125
+ _matrix_backend(left, right, path),
126
+ _matrix_columns(left, right, path),
127
+ _matrix_parameter(left, right, path),
128
+ )
129
+
130
+
131
+ def _matrix_leaf_expression(
132
+ node: Node,
133
+ path: str,
134
+ operation: str,
135
+ /,
136
+ ) -> _MatrixExpression | None:
137
+ if operation == "from_columns":
138
+ return _MatrixExpression(
139
+ _cstr(node.attr("backend")),
140
+ _cstr_seq(node.attr("columns")),
141
+ frozenset({node.args[0].digest}),
142
+ None,
143
+ 0,
144
+ 0,
145
+ True,
146
+ {"op": "input"},
147
+ )
148
+ if operation == "parameter":
149
+ name = _cstr(node.attr("name"))
150
+ if name != "weights":
151
+ errors.raise_compile(
152
+ f"static_inputs.{name}",
153
+ errors.CAPABILITY_MISMATCH,
154
+ "symbolic matrix lowering currently requires the static array"
155
+ " parameter name 'weights'",
156
+ )
157
+ return _MatrixExpression(
158
+ _cstr(node.attr("backend")),
159
+ (),
160
+ frozenset(),
161
+ node,
162
+ 1,
163
+ 0,
164
+ True,
165
+ {"op": "weights"},
166
+ )
167
+ if operation == "literal":
168
+ return _MatrixExpression(
169
+ "",
170
+ (),
171
+ frozenset(),
172
+ None,
173
+ 0,
174
+ 0,
175
+ True,
176
+ {"op": "literal", "value": _matrix_literal(node, path)},
177
+ )
178
+ return None
179
+
180
+
181
+ def _matrix_unary_expression(
182
+ node: Node,
183
+ path: str,
184
+ operation: str,
185
+ /,
186
+ ) -> _MatrixExpression:
187
+ value = _matrix_expression(node.args[0], f"{path}.{operation}.value")
188
+ return _MatrixExpression(
189
+ value.backend,
190
+ value.columns,
191
+ value.source_digests,
192
+ value.parameter,
193
+ value.weights_count,
194
+ value.matmul_count,
195
+ value.matmul_rhs_is_weights,
196
+ {"op": operation, "value": value.tree},
197
+ )
198
+
199
+
200
+ def _matrix_rhs_is_weights(node: Node, operation: str, /) -> bool:
201
+ return operation != "matmul" or (
202
+ node.args[1].op.name == "parameter"
203
+ and _cstr(node.args[1].attr("name")) == "weights"
204
+ )
205
+
206
+
207
+ def _matrix_binary_expression(
208
+ node: Node,
209
+ path: str,
210
+ operation: str,
211
+ /,
212
+ ) -> _MatrixExpression:
213
+ left = _matrix_expression(node.args[0], f"{path}.{operation}.left")
214
+ right = _matrix_expression(node.args[1], f"{path}.{operation}.right")
215
+ backend, columns, parameter = _merge_matrix_expression(left, right, path)
216
+ return _MatrixExpression(
217
+ backend,
218
+ columns,
219
+ left.source_digests | right.source_digests,
220
+ parameter,
221
+ left.weights_count + right.weights_count,
222
+ left.matmul_count + right.matmul_count + (operation == "matmul"),
223
+ left.matmul_rhs_is_weights
224
+ and right.matmul_rhs_is_weights
225
+ and _matrix_rhs_is_weights(node, operation),
226
+ {"left": left.tree, "op": operation, "right": right.tree},
227
+ )
228
+
229
+
230
+ def _matrix_expression(node: Node, path: str, /) -> _MatrixExpression:
231
+ operation = node.op.name
232
+ leaf = _matrix_leaf_expression(node, path, operation)
233
+ if leaf is not None:
234
+ return leaf
235
+ if operation not in _MATRIX_PRIMITIVES:
236
+ _reject_primitive(path, node)
237
+ if operation in ("neg", "not"):
238
+ return _matrix_unary_expression(node, path, operation)
239
+ return _matrix_binary_expression(node, path, operation)
240
+
241
+
242
+ def _static_array_declaration(node: Node, /) -> dict[str, object]:
243
+ dtype = node.attr("dtype")
244
+ shape = node.attr("shape")
245
+ return {
246
+ "backend": _cstr(node.attr("backend")),
247
+ "dtype": dtype.name if isinstance(dtype, CDType) else "",
248
+ "kind": "array",
249
+ "mutability": "static",
250
+ "name": _cstr(node.attr("name")),
251
+ "shape": [
252
+ dimension.value
253
+ for dimension in (shape.items if isinstance(shape, CSeq) else ())
254
+ if isinstance(dimension, CInt)
255
+ ],
256
+ }
257
+
258
+
259
+ def _matrix_program_output(
260
+ program: Program | _LoweringProgram,
261
+ /,
262
+ ) -> tuple[str, Node] | None:
263
+ if len(program.outputs) != 1:
264
+ return None
265
+ output_name, value = program.outputs[0]
266
+ if value._node.op.name != "attach_columns":
267
+ return None
268
+ return output_name, value._node
269
+
270
+
271
+ def _required_matrix_expression(node: Node, output_name: str, /) -> _MatrixExpression:
272
+ path = f"outputs.{output_name}.array"
273
+ matrix = _matrix_expression(node.args[1], path)
274
+ if not matrix.backend or not matrix.columns:
275
+ errors.raise_compile(
276
+ path,
277
+ errors.UNRESOLVED_TYPE,
278
+ "symbolic matrix output requires linalg.from_columns",
279
+ )
280
+ return matrix
281
+
282
+
283
+ def _required_matrix_parameter(
284
+ matrix: _MatrixExpression,
285
+ output_name: str,
286
+ /,
287
+ ) -> Node:
288
+ parameter = matrix.parameter
289
+ if parameter is None:
290
+ errors.raise_compile(
291
+ f"outputs.{output_name}.array",
292
+ errors.UNRESOLVED_TYPE,
293
+ "symbolic matrix output requires one static array parameter",
294
+ )
295
+ if _cstr(parameter.attr("name")) != "weights":
296
+ errors.raise_compile(
297
+ f"outputs.{output_name}.array",
298
+ errors.UNRESOLVED_TYPE,
299
+ "symbolic matrix output requires one static array parameter",
300
+ )
301
+ return parameter
302
+
303
+
304
+ def _require_frozen_matrix_shape(
305
+ matrix: _MatrixExpression,
306
+ output_name: str,
307
+ /,
308
+ ) -> None:
309
+ if (
310
+ matrix.weights_count != 1
311
+ or matrix.matmul_count != 1
312
+ or not matrix.matmul_rhs_is_weights
313
+ ):
314
+ errors.raise_compile(
315
+ f"outputs.{output_name}.array",
316
+ errors.CAPABILITY_MISMATCH,
317
+ "symbolic matrix output requires exactly one matmul whose right"
318
+ " operand is the static 'weights' parameter",
319
+ )
320
+
321
+
322
+ def _attached_matrix_input(
323
+ node: Node,
324
+ matrix: _MatrixExpression,
325
+ output_name: str,
326
+ /,
327
+ ) -> Node:
328
+ attached_input = node.args[0]
329
+ if matrix.source_digests != frozenset({attached_input.digest}):
330
+ errors.raise_compile(
331
+ f"outputs.{output_name}.value",
332
+ errors.SCHEMA_MISMATCH,
333
+ "attached matrix columns must come from the attached table",
334
+ )
335
+ return attached_input
336
+
337
+
338
+ def _matrix_external_node(
339
+ output_name: str,
340
+ node: Node,
341
+ matrix: _MatrixExpression,
342
+ /,
343
+ ) -> dict[str, object]:
344
+ return {
345
+ "id": output_name,
346
+ "input_ports": [
347
+ {"kind": "table", "name": "input", "required": True},
348
+ {"kind": "array", "name": "weights", "required": True},
349
+ ],
350
+ "operator": {
351
+ "kind": "external",
352
+ "name": "symbolic_matrix",
353
+ "options": {
354
+ "columns": list(matrix.columns),
355
+ "expression": matrix.tree,
356
+ "names": list(_cstr_seq(node.attr("names"))),
357
+ },
358
+ "provider": matrix.backend,
359
+ "version": "1",
360
+ },
361
+ "output_ports": [{"kind": "table", "name": "output", "required": True}],
362
+ }
363
+
364
+
365
+ def _matrix_upstream_project(
366
+ program: Program | _LoweringProgram,
367
+ attached_input: Node,
368
+ upstream_id: str,
369
+ mode: str,
370
+ allowed_lateness_micros: int,
371
+ late_policy: str,
372
+ /,
373
+ ) -> dict[str, object]:
374
+ upstream = _LoweringProgram(
375
+ program.name,
376
+ tuple(
377
+ _LoweringValue(item._node)
378
+ for item in program.inputs
379
+ if item._node.op.name == "table_input"
380
+ ),
381
+ ((upstream_id, _LoweringValue(attached_input)),),
382
+ )
383
+ from calc_flow.symbolic.lower.program import _lower_program as _row_local_lower
384
+
385
+ return _row_local_lower(
386
+ upstream,
387
+ mode,
388
+ allowed_lateness_micros,
389
+ late_policy,
390
+ )
391
+
392
+
393
+ def _raise_lowering_invariant(message: str, /) -> Never:
394
+ raise RuntimeError(f"symbolic lowering invariant violated: {message}")
395
+
396
+
397
+ def _project_graph_lists(
398
+ project: dict[str, object],
399
+ /,
400
+ ) -> tuple[list[object], list[object]]:
401
+ graph = project.get("graph")
402
+ if not isinstance(graph, dict):
403
+ _raise_lowering_invariant("project.graph must be a mapping")
404
+ nodes = graph.get("nodes")
405
+ if not isinstance(nodes, list):
406
+ _raise_lowering_invariant("project.graph.nodes must be a list")
407
+ edges = graph.get("edges")
408
+ if not isinstance(edges, list):
409
+ _raise_lowering_invariant("project.graph.edges must be a list")
410
+ return nodes, edges
411
+
412
+
413
+ def _project_document(
414
+ name: str,
415
+ mode: str,
416
+ nodes: list[object],
417
+ edges: list[object],
418
+ /,
419
+ ) -> dict[str, object]:
420
+ project: dict[str, object] = {
421
+ "data_sources": [],
422
+ "format_version": 3,
423
+ "id": name,
424
+ "name": name,
425
+ "runtime": {"mode": mode, "options": {}},
426
+ "graph": {"edges": edges, "name": name, "nodes": nodes},
427
+ }
428
+ project["data_sources"] = [] if mode == "stream" else _data_sources(project)
429
+ return project
430
+
431
+
432
+ def _wire_matrix_node(
433
+ project: dict[str, object],
434
+ external: dict[str, object],
435
+ upstream_id: str,
436
+ output_name: str,
437
+ /,
438
+ ) -> None:
439
+ nodes, edges = _project_graph_lists(project)
440
+ nodes.append(external)
441
+ edges.append(
442
+ {
443
+ "source_node": upstream_id,
444
+ "source_port": "output",
445
+ "target_node": output_name,
446
+ "target_port": "input",
447
+ }
448
+ )
449
+
450
+
451
+ def _lower_matrix_program(
452
+ program: Program | _LoweringProgram,
453
+ mode: str,
454
+ allowed_lateness_micros: int,
455
+ late_policy: str,
456
+ /,
457
+ ) -> dict[str, object] | None:
458
+ output = _matrix_program_output(program)
459
+ if output is None:
460
+ return None
461
+ output_name, node = output
462
+ matrix = _required_matrix_expression(node, output_name)
463
+ parameter = _required_matrix_parameter(matrix, output_name)
464
+ _require_frozen_matrix_shape(matrix, output_name)
465
+ attached_input = _attached_matrix_input(node, matrix, output_name)
466
+ external = _matrix_external_node(output_name, node, matrix)
467
+ upstream_id = f"{output_name}__cf_matrix_input"
468
+ project = _matrix_upstream_project(
469
+ program,
470
+ attached_input,
471
+ upstream_id,
472
+ mode,
473
+ allowed_lateness_micros,
474
+ late_policy,
475
+ )
476
+ _wire_matrix_node(project, external, upstream_id, output_name)
477
+ if mode == "stream":
478
+ project["static_inputs"] = [_static_array_declaration(parameter)]
479
+ else:
480
+ project["data_sources"] = _data_sources(project)
481
+ return project
482
+
483
+
484
+ def _stream_join_nodes(program: Program, /) -> tuple[Node, ...]:
485
+ joins: dict[str, Node] = {}
486
+ for _, value in program.outputs:
487
+ for node in _walk_nodes(value._node):
488
+ if node.op.name == "stream_join":
489
+ joins[node.digest] = node
490
+ return tuple(joins[digest] for digest in sorted(joins))
491
+
492
+
493
+ def _contains_node(root: Node, digest: str, /) -> bool:
494
+ return any(node.digest == digest for node in _walk_nodes(root))
495
+
496
+
497
+ def _contains_primitive(root: Node, primitive: str, /) -> bool:
498
+ return any(node.op.name == primitive for node in _walk_nodes(root))
499
+
500
+
501
+ def _replace_node(root: Node, digest: str, replacement: Node, /) -> Node:
502
+ if root.digest == digest:
503
+ return replacement
504
+ return build(
505
+ root.op.name,
506
+ tuple(_replace_node(child, digest, replacement) for child in root.args),
507
+ dict(root.attrs.entries),
508
+ version=root.op.version,
509
+ )
510
+
511
+
512
+ def _stream_join_node(
513
+ node: Node,
514
+ node_id: str,
515
+ left_schema: tuple[Field, ...],
516
+ right_schema: tuple[Field, ...],
517
+ /,
518
+ ) -> dict[str, object]:
519
+ return {
520
+ "id": node_id,
521
+ "input_ports": [
522
+ {
523
+ "kind": "table",
524
+ "name": "left",
525
+ "required": True,
526
+ "schema": [_field_json(field) for field in left_schema],
527
+ },
528
+ {
529
+ "kind": "table",
530
+ "name": "right",
531
+ "required": True,
532
+ "schema": [_field_json(field) for field in right_schema],
533
+ },
534
+ ],
535
+ "operator": {
536
+ "kind": "stream_join",
537
+ "spec": join_wire_spec(
538
+ JoinSideWire(
539
+ keys=_cstr_seq(node.attr("left_keys")),
540
+ event_time=_cstr(node.attr("left_event_time")),
541
+ prefix=_cstr(node.attr("left_prefix")),
542
+ ),
543
+ JoinSideWire(
544
+ keys=_cstr_seq(node.attr("right_keys")),
545
+ event_time=_cstr(node.attr("right_event_time")),
546
+ prefix=_cstr(node.attr("right_prefix")),
547
+ ),
548
+ bounds_wire(
549
+ _cint(node.attr("before_micros")),
550
+ _cint(node.attr("after_micros")),
551
+ ),
552
+ limits_wire(
553
+ _cint(node.attr("max_state_rows_per_side")),
554
+ _cint(node.attr("max_state_bytes_per_side")),
555
+ _cint(node.attr("max_matches_per_input_batch")),
556
+ ),
557
+ ),
558
+ },
559
+ "output_ports": [],
560
+ }
561
+
562
+
563
+ @dataclass(frozen=True, slots=True)
564
+ class _StreamJoinSide:
565
+ port: str
566
+ node_id: str
567
+ node: Node
568
+ schema: tuple[Field, ...]
569
+
570
+
571
+ @dataclass(frozen=True, slots=True)
572
+ class _StreamJoinPlan:
573
+ node: Node
574
+ node_id: str
575
+ sides: tuple[_StreamJoinSide, _StreamJoinSide]
576
+ output_schema: tuple[Field, ...]
577
+
578
+
579
+ @dataclass(frozen=True, slots=True)
580
+ class _RelationalFragment:
581
+ boundary: Node
582
+ project: dict[str, object] | None
583
+
584
+
585
+ def _connected_input_endpoints(edges: list[object], /) -> set[tuple[str, str]]:
586
+ connected: set[tuple[str, str]] = set()
587
+ for edge in edges:
588
+ if not isinstance(edge, dict):
589
+ continue
590
+ connected.add((str(edge["target_node"]), str(edge.get("target_port", "input"))))
591
+ return connected
592
+
593
+
594
+ def _required_input_port_names(node: dict[str, object], /) -> tuple[str, ...]:
595
+ ports = node.get("input_ports")
596
+ if not isinstance(ports, list):
597
+ return ("input",)
598
+ if not ports:
599
+ return ("input",)
600
+ names: list[str] = []
601
+ for port in ports:
602
+ if not isinstance(port, dict):
603
+ continue
604
+ if port.get("required") is True:
605
+ names.append(str(port["name"]))
606
+ return tuple(names)
607
+
608
+
609
+ def _graph_node(node: object, /) -> dict[str, object]:
610
+ if not isinstance(node, dict):
611
+ _raise_lowering_invariant("graph node must be a mapping")
612
+ return node
613
+
614
+
615
+ def _downstream_input_endpoints(
616
+ nodes: list[object], edges: list[object], /
617
+ ) -> tuple[tuple[str, str], ...]:
618
+ connected = _connected_input_endpoints(edges)
619
+ endpoints: list[tuple[str, str]] = []
620
+ for raw_node in nodes:
621
+ node = _graph_node(raw_node)
622
+ node_id = str(node["id"])
623
+ for name in _required_input_port_names(node):
624
+ endpoint = (node_id, name)
625
+ if endpoint not in connected:
626
+ endpoints.append(endpoint)
627
+ return tuple(sorted(endpoints))
628
+
629
+
630
+ def _pin_table_output(
631
+ nodes: list[object], node_id: str, schema: tuple[Field, ...], /
632
+ ) -> None:
633
+ for node in nodes:
634
+ if isinstance(node, dict) and node.get("id") == node_id:
635
+ node["output_ports"] = [
636
+ {
637
+ "kind": "table",
638
+ "name": "output",
639
+ "required": True,
640
+ "schema": [_field_json(field) for field in schema],
641
+ }
642
+ ]
643
+ return
644
+ _raise_lowering_invariant(f"missing stream join input stage {node_id!r}")
645
+
646
+
647
+ def _required_stream_join(program: Program, joins: tuple[Node, ...], /) -> Node:
648
+ if len(joins) != 1:
649
+ errors.raise_compile(
650
+ program.name,
651
+ errors.CAPABILITY_MISMATCH,
652
+ "SCE-17 supports exactly one unique symbolic stream join per program",
653
+ )
654
+ return joins[0]
655
+
656
+
657
+ def _check_stream_join_outputs(program: Program, join: Node, /) -> None:
658
+ for output_name, value in program.outputs:
659
+ if _contains_node(value._node, join.digest):
660
+ continue
661
+ errors.raise_compile(
662
+ f"outputs.{output_name}",
663
+ errors.CAPABILITY_MISMATCH,
664
+ "every output in a symbolic stream-join program must descend"
665
+ " from its one shared join",
666
+ )
667
+
668
+
669
+ def _check_stream_join_inputs(program: Program, join: Node, /) -> None:
670
+ for side_name, side in zip(("left", "right"), join.args, strict=True):
671
+ if not _contains_primitive(side, "attach_columns"):
672
+ continue
673
+ errors.raise_compile(
674
+ f"{program.name}.stream_join.{side_name}",
675
+ errors.CAPABILITY_MISMATCH,
676
+ "matrix attachment around a symbolic stream join is not supported",
677
+ )
678
+
679
+
680
+ def _stream_join_plan(
681
+ program: Program, analyzer: _Analyzer, join: Node, /
682
+ ) -> _StreamJoinPlan:
683
+ left_facts = analyzer.table(join.args[0], f"{program.name}.stream_join.left")
684
+ right_facts = analyzer.table(join.args[1], f"{program.name}.stream_join.right")
685
+ join_facts = analyzer.table(join, f"{program.name}.stream_join")
686
+ join_id = f"cf_stream_join_{join.digest[:16]}"
687
+ return _StreamJoinPlan(
688
+ join,
689
+ join_id,
690
+ (
691
+ _StreamJoinSide(
692
+ "left",
693
+ f"{join_id}__left",
694
+ join.args[0],
695
+ left_facts.schema,
696
+ ),
697
+ _StreamJoinSide(
698
+ "right",
699
+ f"{join_id}__right",
700
+ join.args[1],
701
+ right_facts.schema,
702
+ ),
703
+ ),
704
+ join_facts.schema,
705
+ )
706
+
707
+
708
+ def _stream_join_upstream_outputs(
709
+ plan: _StreamJoinPlan, /
710
+ ) -> tuple[tuple[str, _LoweringValue], ...]:
711
+ return tuple(
712
+ (side.node_id, _LoweringValue(side.node))
713
+ for side in plan.sides
714
+ if side.node.op.name != "table_input"
715
+ )
716
+
717
+
718
+ def _input_reaches_join_side(
719
+ input_node: Node, upstream_sides: tuple[Node, ...], /
720
+ ) -> bool:
721
+ return any(_contains_node(side, input_node.digest) for side in upstream_sides)
722
+
723
+
724
+ def _stream_join_upstream_inputs(
725
+ program: Program,
726
+ upstream_outputs: tuple[tuple[str, _LoweringValue], ...],
727
+ /,
728
+ ) -> tuple[_LoweringValue, ...]:
729
+ upstream_sides = tuple(value._node for _, value in upstream_outputs)
730
+ inputs: list[_LoweringValue] = []
731
+ for value in program.inputs:
732
+ if value._node.op.name != "table_input":
733
+ continue
734
+ if _input_reaches_join_side(value._node, upstream_sides):
735
+ inputs.append(_LoweringValue(value._node))
736
+ return tuple(inputs)
737
+
738
+
739
+ def _stream_join_upstream_project(
740
+ program: Program,
741
+ plan: _StreamJoinPlan,
742
+ mode: str,
743
+ allowed_lateness_micros: int,
744
+ late_policy: str,
745
+ /,
746
+ ) -> dict[str, object]:
747
+ outputs = _stream_join_upstream_outputs(plan)
748
+ if not outputs:
749
+ return _project_document(program.name, mode, [], [])
750
+ inputs = _stream_join_upstream_inputs(program, outputs)
751
+ from calc_flow.symbolic.lower.program import _lower_program as _row_local_lower
752
+
753
+ return _row_local_lower(
754
+ _LoweringProgram(program.name, inputs, outputs),
755
+ mode,
756
+ allowed_lateness_micros,
757
+ late_policy,
758
+ )
759
+
760
+
761
+ def _wire_stream_join_inputs(
762
+ nodes: list[object], edges: list[object], plan: _StreamJoinPlan, /
763
+ ) -> None:
764
+ for side in plan.sides:
765
+ if side.node.op.name == "table_input":
766
+ continue
767
+ _pin_table_output(nodes, side.node_id, side.schema)
768
+ edges.append(
769
+ {
770
+ "source_node": side.node_id,
771
+ "source_port": "output",
772
+ "target_node": plan.node_id,
773
+ "target_port": side.port,
774
+ }
775
+ )
776
+
777
+
778
+ def _direct_stream_join_output(program: Program, plan: _StreamJoinPlan, /) -> bool:
779
+ if len(program.outputs) != 1:
780
+ return False
781
+ return program.outputs[0][1]._node.digest == plan.node.digest
782
+
783
+
784
+ def _wire_stream_join_downstream(
785
+ program: Program,
786
+ plan: _StreamJoinPlan,
787
+ mode: str,
788
+ allowed_lateness_micros: int,
789
+ late_policy: str,
790
+ upstream_nodes: list[object],
791
+ upstream_edges: list[object],
792
+ /,
793
+ ) -> None:
794
+ from calc_flow.symbolic.expr import table_input
795
+
796
+ virtual_id = f"{plan.node_id}__output"
797
+ virtual = table_input(virtual_id, schema=plan.output_schema)
798
+ downstream = _LoweringProgram(
799
+ program.name,
800
+ (_LoweringValue(virtual._node),),
801
+ tuple(
802
+ (
803
+ output_name,
804
+ _LoweringValue(
805
+ _replace_node(value._node, plan.node.digest, virtual._node)
806
+ ),
807
+ )
808
+ for output_name, value in program.outputs
809
+ ),
810
+ )
811
+ from calc_flow.symbolic.lower.program import _lower_program as _row_local_lower
812
+
813
+ downstream_project = _row_local_lower(
814
+ downstream,
815
+ mode,
816
+ allowed_lateness_micros,
817
+ late_policy,
818
+ )
819
+ downstream_nodes, downstream_edges = _project_graph_lists(downstream_project)
820
+ endpoints = _downstream_input_endpoints(downstream_nodes, downstream_edges)
821
+ if not endpoints:
822
+ _raise_lowering_invariant("stream join downstream has no input endpoint")
823
+ upstream_nodes.extend(downstream_nodes)
824
+ upstream_edges.extend(downstream_edges)
825
+ upstream_edges.extend(
826
+ {
827
+ "source_node": plan.node_id,
828
+ "source_port": "output",
829
+ "target_node": target_node,
830
+ "target_port": target_port,
831
+ }
832
+ for target_node, target_port in endpoints
833
+ )
834
+
835
+
836
+ def _lower_stream_join_program(
837
+ program: Program,
838
+ analyzer: _Analyzer,
839
+ mode: str,
840
+ allowed_lateness_micros: int,
841
+ late_policy: str,
842
+ /,
843
+ ) -> dict[str, object] | None:
844
+ joins = _stream_join_nodes(program)
845
+ if not joins:
846
+ return None
847
+ if _requires_relational_dag_lowering(program, joins):
848
+ return _lower_relational_dag_program(
849
+ program,
850
+ analyzer,
851
+ mode,
852
+ allowed_lateness_micros,
853
+ late_policy,
854
+ joins,
855
+ )
856
+ join = _required_stream_join(program, joins)
857
+ _check_stream_join_outputs(program, join)
858
+ _check_stream_join_inputs(program, join)
859
+ plan = _stream_join_plan(program, analyzer, join)
860
+ upstream_project = _stream_join_upstream_project(
861
+ program,
862
+ plan,
863
+ mode,
864
+ allowed_lateness_micros,
865
+ late_policy,
866
+ )
867
+ upstream_nodes, upstream_edges = _project_graph_lists(upstream_project)
868
+ _wire_stream_join_inputs(upstream_nodes, upstream_edges, plan)
869
+ upstream_nodes.append(
870
+ _stream_join_node(
871
+ plan.node,
872
+ plan.node_id,
873
+ plan.sides[0].schema,
874
+ plan.sides[1].schema,
875
+ )
876
+ )
877
+ if _direct_stream_join_output(program, plan):
878
+ return upstream_project
879
+ _wire_stream_join_downstream(
880
+ program,
881
+ plan,
882
+ mode,
883
+ allowed_lateness_micros,
884
+ late_policy,
885
+ upstream_nodes,
886
+ upstream_edges,
887
+ )
888
+ return upstream_project
889
+
890
+
891
+ def _requires_relational_dag_lowering(
892
+ program: Program, joins: tuple[Node, ...], /
893
+ ) -> bool:
894
+ if len(joins) != 1 or joins[0].op.version >= 2:
895
+ return True
896
+ join = joins[0]
897
+ return any(
898
+ not _contains_node(value._node, join.digest) for _, value in program.outputs
899
+ )
900
+
901
+
902
+ def _relational_boundary(node: Node, path: str, /) -> Node:
903
+ current = node
904
+ while current.op.name in ("project", "filter", "with_columns"):
905
+ current = current.args[0]
906
+ if current.op.name in ("table_input", "stream_join"):
907
+ return current
908
+ if _contains_primitive(current, "attach_columns"):
909
+ errors.raise_compile(
910
+ path,
911
+ errors.CAPABILITY_MISMATCH,
912
+ "matrix attachment around a symbolic stream join is not supported",
913
+ )
914
+ _reject_primitive(path, current)
915
+
916
+
917
+ def _virtual_relational_input(node_id: str, facts: TableFacts, /) -> Node:
918
+ from calc_flow.symbolic.expr import table_input
919
+
920
+ return table_input(
921
+ node_id,
922
+ schema=facts.schema,
923
+ entity_by=facts.entity_by,
924
+ event_time=facts.event_time,
925
+ sequence_by=facts.sequence_by,
926
+ )._node
927
+
928
+
929
+ def _reachable_relational_sources(program: Program, /) -> frozenset[str]:
930
+ return frozenset(
931
+ node.digest
932
+ for _, value in program.outputs
933
+ for node in _walk_nodes(value._node)
934
+ if node.op.name == "table_input"
935
+ )
936
+
937
+
938
+ def _relational_source_name(node: Node, reserved_ids: frozenset[str], /) -> str:
939
+ declared_name = _cstr(node.attr("name"))
940
+ if declared_name is None:
941
+ _raise_lowering_invariant("relational source is missing its declared name")
942
+ if declared_name in reserved_ids:
943
+ return f"cf_source_{node.digest[:16]}"
944
+ return declared_name
945
+
946
+
947
+ def _relational_source_nodes(
948
+ program: Program, reserved_ids: frozenset[str], /
949
+ ) -> tuple[dict[str, str], list[dict[str, object]]]:
950
+ reachable = _reachable_relational_sources(program)
951
+ by_digest: dict[str, str] = {}
952
+ nodes: list[dict[str, object]] = []
953
+ for value in program.inputs:
954
+ node = value._node
955
+ if node.op.name != "table_input" or node.digest not in reachable:
956
+ continue
957
+ name = _relational_source_name(node, reserved_ids)
958
+ schema = _schema_fields(node.attr("schema"))
959
+ by_digest[node.digest] = name
960
+ nodes.append(
961
+ _expression_node(
962
+ name,
963
+ [_quote_identifier(field.name) for field in schema],
964
+ None,
965
+ schema,
966
+ schema,
967
+ )
968
+ )
969
+ return by_digest, nodes
970
+
971
+
972
+ def _relational_upstream_id(
973
+ boundary: Node,
974
+ sources: dict[str, str],
975
+ joins: dict[str, _StreamJoinPlan],
976
+ path: str,
977
+ /,
978
+ ) -> str:
979
+ if boundary.op.name == "table_input":
980
+ source_id = sources.get(boundary.digest)
981
+ if source_id is None:
982
+ _raise_lowering_invariant(
983
+ f"missing declared source for relational boundary at {path}"
984
+ )
985
+ return source_id
986
+ plan = joins.get(boundary.digest)
987
+ if plan is None:
988
+ _raise_lowering_invariant(
989
+ f"missing physical join for relational boundary at {path}"
990
+ )
991
+ return plan.node_id
992
+
993
+
994
+ def _relational_fragment(
995
+ program: Program,
996
+ analyzer: _Analyzer,
997
+ expression: Node,
998
+ boundary: Node,
999
+ output_id: str,
1000
+ path: str,
1001
+ mode: str,
1002
+ allowed_lateness_micros: int,
1003
+ late_policy: str,
1004
+ /,
1005
+ ) -> dict[str, object]:
1006
+ if boundary.op.name == "stream_join":
1007
+ facts = analyzer.table(boundary, f"{path}.boundary")
1008
+ declared = _virtual_relational_input(
1009
+ f"cf_join_output_{boundary.digest[:16]}", facts
1010
+ )
1011
+ lowered = _replace_node(expression, boundary.digest, declared)
1012
+ else:
1013
+ declared = boundary
1014
+ lowered = expression
1015
+ from calc_flow.symbolic.lower.program import _lower_program as _row_local_lower
1016
+
1017
+ return _row_local_lower(
1018
+ _LoweringProgram(
1019
+ program.name,
1020
+ (_LoweringValue(declared),),
1021
+ ((output_id, _LoweringValue(lowered)),),
1022
+ ),
1023
+ mode,
1024
+ allowed_lateness_micros,
1025
+ late_policy,
1026
+ )
1027
+
1028
+
1029
+ def _relational_output_fragments(
1030
+ program: Program,
1031
+ analyzer: _Analyzer,
1032
+ mode: str,
1033
+ allowed_lateness_micros: int,
1034
+ late_policy: str,
1035
+ /,
1036
+ ) -> dict[str, _RelationalFragment]:
1037
+ fragments: dict[str, _RelationalFragment] = {}
1038
+ for output_name, value in program.outputs:
1039
+ path = f"outputs.{output_name}"
1040
+ boundary = _relational_boundary(value._node, path)
1041
+ fragments[output_name] = _RelationalFragment(
1042
+ boundary,
1043
+ _relational_fragment(
1044
+ program,
1045
+ analyzer,
1046
+ value._node,
1047
+ boundary,
1048
+ output_name,
1049
+ path,
1050
+ mode,
1051
+ allowed_lateness_micros,
1052
+ late_policy,
1053
+ ),
1054
+ )
1055
+ return fragments
1056
+
1057
+
1058
+ def _relational_join_fragments(
1059
+ program: Program,
1060
+ analyzer: _Analyzer,
1061
+ plans: dict[str, _StreamJoinPlan],
1062
+ mode: str,
1063
+ allowed_lateness_micros: int,
1064
+ late_policy: str,
1065
+ /,
1066
+ ) -> dict[tuple[str, str], _RelationalFragment]:
1067
+ fragments: dict[tuple[str, str], _RelationalFragment] = {}
1068
+ for plan in plans.values():
1069
+ for side in plan.sides:
1070
+ path = f"{program.name}.{plan.node_id}.{side.port}"
1071
+ boundary = _relational_boundary(side.node, path)
1072
+ project = None
1073
+ if side.node.digest != boundary.digest:
1074
+ project = _relational_fragment(
1075
+ program,
1076
+ analyzer,
1077
+ side.node,
1078
+ boundary,
1079
+ side.node_id,
1080
+ path,
1081
+ mode,
1082
+ allowed_lateness_micros,
1083
+ late_policy,
1084
+ )
1085
+ fragments[(plan.node.digest, side.port)] = _RelationalFragment(
1086
+ boundary,
1087
+ project,
1088
+ )
1089
+ return fragments
1090
+
1091
+
1092
+ def _wire_relational_fragment(
1093
+ nodes: list[object],
1094
+ edges: list[object],
1095
+ fragment: dict[str, object],
1096
+ upstream_id: str,
1097
+ output_id: str,
1098
+ output_schema: tuple[Field, ...] | None,
1099
+ /,
1100
+ ) -> None:
1101
+ fragment_nodes, fragment_edges = _project_graph_lists(fragment)
1102
+ endpoints = _downstream_input_endpoints(fragment_nodes, fragment_edges)
1103
+ if len(endpoints) != 1:
1104
+ _raise_lowering_invariant(
1105
+ f"relational fragment {output_id!r} has {len(endpoints)} input endpoints"
1106
+ )
1107
+ if output_schema is not None:
1108
+ _pin_table_output(fragment_nodes, output_id, output_schema)
1109
+ nodes.extend(fragment_nodes)
1110
+ edges.extend(fragment_edges)
1111
+ target_node, target_port = endpoints[0]
1112
+ edges.append(
1113
+ {
1114
+ "source_node": upstream_id,
1115
+ "source_port": "output",
1116
+ "target_node": target_node,
1117
+ "target_port": target_port,
1118
+ }
1119
+ )
1120
+
1121
+
1122
+ def _wire_relational_join_side(
1123
+ program: Program,
1124
+ plan: _StreamJoinPlan,
1125
+ side: _StreamJoinSide,
1126
+ fragment: _RelationalFragment,
1127
+ sources: dict[str, str],
1128
+ joins: dict[str, _StreamJoinPlan],
1129
+ nodes: list[object],
1130
+ edges: list[object],
1131
+ /,
1132
+ ) -> None:
1133
+ path = f"{program.name}.{plan.node_id}.{side.port}"
1134
+ upstream_id = _relational_upstream_id(fragment.boundary, sources, joins, path)
1135
+ if fragment.project is None:
1136
+ edges.append(
1137
+ {
1138
+ "source_node": upstream_id,
1139
+ "source_port": "output",
1140
+ "target_node": plan.node_id,
1141
+ "target_port": side.port,
1142
+ }
1143
+ )
1144
+ return
1145
+ _wire_relational_fragment(
1146
+ nodes,
1147
+ edges,
1148
+ fragment.project,
1149
+ upstream_id,
1150
+ side.node_id,
1151
+ side.schema,
1152
+ )
1153
+ edges.append(
1154
+ {
1155
+ "source_node": side.node_id,
1156
+ "source_port": "output",
1157
+ "target_node": plan.node_id,
1158
+ "target_port": side.port,
1159
+ }
1160
+ )
1161
+
1162
+
1163
+ def _wire_relational_output(
1164
+ output_name: str,
1165
+ fragment: _RelationalFragment,
1166
+ sources: dict[str, str],
1167
+ joins: dict[str, _StreamJoinPlan],
1168
+ nodes: list[object],
1169
+ edges: list[object],
1170
+ /,
1171
+ ) -> None:
1172
+ path = f"outputs.{output_name}"
1173
+ upstream_id = _relational_upstream_id(fragment.boundary, sources, joins, path)
1174
+ if fragment.project is None:
1175
+ _raise_lowering_invariant(f"relational output {output_name!r} has no fragment")
1176
+ _wire_relational_fragment(
1177
+ nodes,
1178
+ edges,
1179
+ fragment.project,
1180
+ upstream_id,
1181
+ output_name,
1182
+ None,
1183
+ )
1184
+
1185
+
1186
+ def _relational_join_plans(
1187
+ program: Program,
1188
+ analyzer: _Analyzer,
1189
+ join_nodes: tuple[Node, ...],
1190
+ /,
1191
+ ) -> dict[str, _StreamJoinPlan]:
1192
+ for join in join_nodes:
1193
+ _check_stream_join_inputs(program, join)
1194
+ return {
1195
+ join.digest: _stream_join_plan(program, analyzer, join) for join in join_nodes
1196
+ }
1197
+
1198
+
1199
+ def _relational_reserved_ids(
1200
+ program: Program,
1201
+ plans: dict[str, _StreamJoinPlan],
1202
+ join_fragments: dict[tuple[str, str], _RelationalFragment],
1203
+ output_fragments: dict[str, _RelationalFragment],
1204
+ /,
1205
+ ) -> frozenset[str]:
1206
+ reserved = {
1207
+ *(output_name for output_name, _ in program.outputs),
1208
+ *(plan.node_id for plan in plans.values()),
1209
+ *(side.node_id for plan in plans.values() for side in plan.sides),
1210
+ }
1211
+ for fragment in (*join_fragments.values(), *output_fragments.values()):
1212
+ if fragment.project is None:
1213
+ continue
1214
+ nodes, _ = _project_graph_lists(fragment.project)
1215
+ reserved.update(str(_graph_node(node)["id"]) for node in nodes)
1216
+ return frozenset(reserved)
1217
+
1218
+
1219
+ def _append_relational_joins(
1220
+ program: Program,
1221
+ join_nodes: tuple[Node, ...],
1222
+ sources: dict[str, str],
1223
+ plans: dict[str, _StreamJoinPlan],
1224
+ fragments: dict[tuple[str, str], _RelationalFragment],
1225
+ nodes: list[object],
1226
+ edges: list[object],
1227
+ /,
1228
+ ) -> None:
1229
+ for join in join_nodes:
1230
+ plan = plans[join.digest]
1231
+ nodes.append(
1232
+ _stream_join_node(
1233
+ plan.node,
1234
+ plan.node_id,
1235
+ plan.sides[0].schema,
1236
+ plan.sides[1].schema,
1237
+ )
1238
+ )
1239
+ for side in plan.sides:
1240
+ _wire_relational_join_side(
1241
+ program,
1242
+ plan,
1243
+ side,
1244
+ fragments[(plan.node.digest, side.port)],
1245
+ sources,
1246
+ plans,
1247
+ nodes,
1248
+ edges,
1249
+ )
1250
+
1251
+
1252
+ def _append_relational_outputs(
1253
+ program: Program,
1254
+ sources: dict[str, str],
1255
+ plans: dict[str, _StreamJoinPlan],
1256
+ fragments: dict[str, _RelationalFragment],
1257
+ nodes: list[object],
1258
+ edges: list[object],
1259
+ /,
1260
+ ) -> None:
1261
+ for output_name, _ in program.outputs:
1262
+ _wire_relational_output(
1263
+ output_name,
1264
+ fragments[output_name],
1265
+ sources,
1266
+ plans,
1267
+ nodes,
1268
+ edges,
1269
+ )
1270
+
1271
+
1272
+ def _lower_relational_dag_program(
1273
+ program: Program,
1274
+ analyzer: _Analyzer,
1275
+ mode: str,
1276
+ allowed_lateness_micros: int,
1277
+ late_policy: str,
1278
+ join_nodes: tuple[Node, ...],
1279
+ /,
1280
+ ) -> dict[str, object]:
1281
+ plans = _relational_join_plans(program, analyzer, join_nodes)
1282
+ join_fragments = _relational_join_fragments(
1283
+ program,
1284
+ analyzer,
1285
+ plans,
1286
+ mode,
1287
+ allowed_lateness_micros,
1288
+ late_policy,
1289
+ )
1290
+ output_fragments = _relational_output_fragments(
1291
+ program,
1292
+ analyzer,
1293
+ mode,
1294
+ allowed_lateness_micros,
1295
+ late_policy,
1296
+ )
1297
+ reserved_ids = _relational_reserved_ids(
1298
+ program,
1299
+ plans,
1300
+ join_fragments,
1301
+ output_fragments,
1302
+ )
1303
+ sources, source_nodes = _relational_source_nodes(program, reserved_ids)
1304
+ nodes: list[object] = list(source_nodes)
1305
+ edges: list[object] = []
1306
+ _append_relational_joins(
1307
+ program,
1308
+ join_nodes,
1309
+ sources,
1310
+ plans,
1311
+ join_fragments,
1312
+ nodes,
1313
+ edges,
1314
+ )
1315
+ _append_relational_outputs(
1316
+ program,
1317
+ sources,
1318
+ plans,
1319
+ output_fragments,
1320
+ nodes,
1321
+ edges,
1322
+ )
1323
+ typed_nodes = [_graph_node(node) for node in nodes]
1324
+ typed_edges = [_graph_node(edge) for edge in edges]
1325
+ typed_nodes, typed_edges = _deduplicate_node_ids(typed_nodes, typed_edges)
1326
+ return _project_document(program.name, mode, typed_nodes, typed_edges)
1327
+
1328
+
1329
+ def _required_segment_state_plan(
1330
+ rolling: _RollingPipeline | None, cross: _CrossSectionPlan | None, /
1331
+ ) -> _RollingPipeline | _CrossSectionPlan:
1332
+ plan = rolling if rolling is not None else cross
1333
+ if plan is None:
1334
+ raise RuntimeError("missing state plan for a stateful symbolic segment")
1335
+ return plan
1336
+
1337
+
1338
+ def _walk_nodes(root: Node, /) -> tuple[Node, ...]:
1339
+ nodes: list[Node] = []
1340
+
1341
+ def visit(node: Node) -> None:
1342
+ nodes.append(node)
1343
+ for child in node.args:
1344
+ visit(child)
1345
+
1346
+ visit(root)
1347
+ return tuple(nodes)
1348
+
1349
+
1350
+ def _deduplicated_expression_signature(
1351
+ node: dict[str, object],
1352
+ incoming: dict[str, list[dict[str, object]]],
1353
+ aliases: dict[str, str],
1354
+ output_ids: frozenset[str],
1355
+ /,
1356
+ ) -> str | None:
1357
+ node_id = str(node["id"])
1358
+ operator = node["operator"]
1359
+ if not isinstance(operator, dict):
1360
+ return None
1361
+ if operator.get("kind") != "expression" or node_id in output_ids:
1362
+ return None
1363
+ node_incoming = incoming.get(node_id, [])
1364
+ if len(node_incoming) != 1:
1365
+ return None
1366
+ edge = node_incoming[0]
1367
+ source_node = str(edge["source_node"])
1368
+ source = aliases.get(source_node, source_node)
1369
+ return _canonical(
1370
+ {
1371
+ "input_ports": node.get("input_ports", []),
1372
+ "operator": operator,
1373
+ "output_ports": node.get("output_ports", []),
1374
+ "source_node": source,
1375
+ "source_port": edge.get("source_port", "output"),
1376
+ "target_port": edge.get("target_port", "input"),
1377
+ }
1378
+ )
1379
+
1380
+
1381
+ def _resolve_node_alias(node_id: str, aliases: dict[str, str], /) -> str:
1382
+ while node_id in aliases:
1383
+ node_id = aliases[node_id]
1384
+ return node_id
1385
+
1386
+
1387
+ def _rewrite_deduplicated_edges(
1388
+ edges: list[dict[str, object]], aliases: dict[str, str], /
1389
+ ) -> list[dict[str, object]]:
1390
+ rewritten: list[dict[str, object]] = []
1391
+ seen: set[tuple[str, str, str, str]] = set()
1392
+ for edge in edges:
1393
+ target = str(edge["target_node"])
1394
+ if target in aliases:
1395
+ continue
1396
+ source = _resolve_node_alias(str(edge["source_node"]), aliases)
1397
+ normalized = {**edge, "source_node": source, "target_node": target}
1398
+ identity = (
1399
+ source,
1400
+ str(normalized.get("source_port", "output")),
1401
+ target,
1402
+ str(normalized.get("target_port", "input")),
1403
+ )
1404
+ if identity in seen:
1405
+ continue
1406
+ seen.add(identity)
1407
+ rewritten.append(normalized)
1408
+ return rewritten
1409
+
1410
+
1411
+ def _deduplicate_pure_expression_nodes(
1412
+ nodes: list[dict[str, object]],
1413
+ edges: list[dict[str, object]],
1414
+ output_ids: frozenset[str],
1415
+ /,
1416
+ ) -> tuple[list[dict[str, object]], list[dict[str, object]]]:
1417
+ """Share identical connected expression stages across output branches."""
1418
+
1419
+ incoming: dict[str, list[dict[str, object]]] = {}
1420
+ for edge in edges:
1421
+ incoming.setdefault(str(edge["target_node"]), []).append(edge)
1422
+ aliases: dict[str, str] = {}
1423
+ signatures: dict[str, str] = {}
1424
+ kept: list[dict[str, object]] = []
1425
+ for node in nodes:
1426
+ node_id = str(node["id"])
1427
+ signature = _deduplicated_expression_signature(
1428
+ node, incoming, aliases, output_ids
1429
+ )
1430
+ if signature is None:
1431
+ kept.append(node)
1432
+ continue
1433
+ canonical = signatures.get(signature)
1434
+ if canonical is None:
1435
+ signatures[signature] = node_id
1436
+ kept.append(node)
1437
+ else:
1438
+ aliases[node_id] = canonical
1439
+ return kept, _rewrite_deduplicated_edges(edges, aliases)
1440
+
1441
+
1442
+ def _deduplicate_node_ids(
1443
+ nodes: list[dict[str, object]], edges: list[dict[str, object]], /
1444
+ ) -> tuple[list[dict[str, object]], list[dict[str, object]]]:
1445
+ """Collapse repeated references to one already-shared physical stage."""
1446
+
1447
+ by_id: dict[str, str] = {}
1448
+ unique_nodes: list[dict[str, object]] = []
1449
+ for node in nodes:
1450
+ node_id = str(node["id"])
1451
+ canonical = _canonical(node)
1452
+ existing = by_id.get(node_id)
1453
+ if existing is None:
1454
+ by_id[node_id] = canonical
1455
+ unique_nodes.append(node)
1456
+ elif existing != canonical:
1457
+ raise RuntimeError(
1458
+ f"symbolic optimizer emitted conflicting node id {node_id!r}"
1459
+ )
1460
+ unique_edges: list[dict[str, object]] = []
1461
+ seen: set[tuple[str, str, str, str]] = set()
1462
+ for edge in edges:
1463
+ identity = (
1464
+ str(edge["source_node"]),
1465
+ str(edge.get("source_port", "output")),
1466
+ str(edge["target_node"]),
1467
+ str(edge.get("target_port", "input")),
1468
+ )
1469
+ if identity not in seen:
1470
+ seen.add(identity)
1471
+ unique_edges.append(edge)
1472
+ return unique_nodes, unique_edges