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,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
|