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,840 @@
1
+ """The row-local program orchestrator and compile entry points.
2
+
3
+ Deterministic lowering of symbolic programs of symbolic programs to strict project-v3.
4
+
5
+ The lowerer resolves each declared table output into one fused row-local
6
+ segment, renders the segment as DataFusion SQL inside strict project-v3
7
+ ``expression`` nodes, and hands the document to the existing Rust graph
8
+ compiler for final port, schema, topology, and fingerprint validation. No data
9
+ object, source, sink, or runner is accepted here, and no symbolic Python runs
10
+ while a compiled plan executes.
11
+ """
12
+
13
+ from __future__ import annotations
14
+
15
+ from calc_flow.capabilities import RuntimeCapabilities
16
+ from calc_flow.pipeline import (
17
+ BatchExecutionPlan,
18
+ Runtime,
19
+ StreamExecutionPlan,
20
+ StreamRequirements,
21
+ _canonical,
22
+ )
23
+ from calc_flow.symbolic import errors
24
+ from calc_flow.symbolic.analyzer import (
25
+ _Analyzer,
26
+ _require_mode,
27
+ _run,
28
+ _schema_fields,
29
+ )
30
+ from calc_flow.symbolic.domains import type_name
31
+ from calc_flow.symbolic.lower.planners import (
32
+ _check_declared_inputs,
33
+ _CrossSectionPlan,
34
+ _LoweringProgram,
35
+ _plan_cross_section,
36
+ _plan_rolling,
37
+ _share_cross_section_plans,
38
+ _share_rolling_plans,
39
+ )
40
+ from calc_flow.symbolic.lower.segments import (
41
+ _TABLE_OUTPUT_PRIMITIVES,
42
+ _U64_MAX,
43
+ _CompileCacheKey,
44
+ _cstr,
45
+ _expression_node,
46
+ _find_cross_section,
47
+ _find_rolling,
48
+ _quote_identifier,
49
+ _reject_primitive,
50
+ _resolve_table,
51
+ _Segment,
52
+ _select_item,
53
+ _sql,
54
+ )
55
+ from calc_flow.symbolic.lower.strategies import (
56
+ _deduplicate_node_ids,
57
+ _deduplicate_pure_expression_nodes,
58
+ _lower_matrix_program,
59
+ _lower_stream_join_program,
60
+ _project_document,
61
+ _required_segment_state_plan,
62
+ _stream_join_nodes,
63
+ )
64
+ from calc_flow.symbolic.nodes import (
65
+ build,
66
+ )
67
+ from calc_flow.symbolic.optimizer import expression_refs, extract_common
68
+ from calc_flow.symbolic.program import Program
69
+
70
+
71
+ # The lowerer keeps per-output segment staging in one deterministic pass:
72
+ # stage order, edge wiring, and id assignment are semantic, so the rolling,
73
+ # prefilter, and CSE stages stay in one place.
74
+ def _lower_program(
75
+ program: Program | _LoweringProgram,
76
+ mode: str,
77
+ allowed_lateness_micros: int,
78
+ late_policy: str,
79
+ /,
80
+ ) -> dict[str, object]:
81
+ # #lizard forgives
82
+ matrix_project = _lower_matrix_program(
83
+ program,
84
+ mode,
85
+ allowed_lateness_micros,
86
+ late_policy,
87
+ )
88
+ if matrix_project is not None:
89
+ return matrix_project
90
+ _check_declared_inputs(program)
91
+ segments = []
92
+ for output_name, value in program.outputs:
93
+ node = value._node
94
+ path = f"outputs.{output_name}"
95
+ if node.op.name not in _TABLE_OUTPUT_PRIMITIVES:
96
+ _reject_primitive(path, node)
97
+ segments.append((output_name, _resolve_table(node, path)))
98
+ consumed = {segment.input_node.digest for _, segment in segments}
99
+ multi_output_lineages = {
100
+ digest
101
+ for digest in consumed
102
+ if sum(1 for _, segment in segments if segment.input_node.digest == digest) > 1
103
+ }
104
+ fanout = len(program.inputs) > 1 or bool(multi_output_lineages)
105
+ plans = {
106
+ output_name: _plan_rolling(
107
+ output_name,
108
+ segment,
109
+ f"outputs.{output_name}",
110
+ allowed_lateness_micros,
111
+ late_policy,
112
+ )
113
+ for output_name, segment in segments
114
+ }
115
+ plans = _share_rolling_plans(segments, plans)
116
+ rolling_digests = {
117
+ segment.input_node.digest
118
+ for (output_name, segment), plan in zip(segments, plans.values(), strict=True)
119
+ if plan is not None
120
+ }
121
+ cross_plans: dict[str, _CrossSectionPlan | None] = {}
122
+ cross_segments: list[tuple[str, _Segment]] = []
123
+ for (output_name, segment), rolling in zip(segments, plans.values(), strict=True):
124
+ if rolling is not None:
125
+ # Cross-section planning runs over the rolling-rewritten
126
+ # environment so a measured value may be a materialized rolling
127
+ # output column.
128
+ after_rolling = _Segment(
129
+ segment.input_node,
130
+ segment.fields,
131
+ rolling.env,
132
+ None,
133
+ rolling.post_predicate,
134
+ )
135
+ cross_segments.append((output_name, after_rolling))
136
+ cross_plans[output_name] = _plan_cross_section(
137
+ output_name,
138
+ after_rolling,
139
+ f"outputs.{output_name}",
140
+ allowed_lateness_micros,
141
+ late_policy,
142
+ rolling.output_fields,
143
+ )
144
+ else:
145
+ cross_segments.append((output_name, segment))
146
+ cross_plans[output_name] = _plan_cross_section(
147
+ output_name,
148
+ segment,
149
+ f"outputs.{output_name}",
150
+ allowed_lateness_micros,
151
+ late_policy,
152
+ None,
153
+ )
154
+ cross_plans = _share_cross_section_plans(
155
+ cross_segments,
156
+ {
157
+ output_name: None
158
+ if plans[output_name] is None
159
+ else plans[output_name].node_id
160
+ for output_name, _ in segments
161
+ },
162
+ cross_plans,
163
+ )
164
+ direct_cross_section_digests = {
165
+ segment.input_node.digest
166
+ for (output_name, segment), rolling, cross in zip(
167
+ segments, plans.values(), cross_plans.values(), strict=True
168
+ )
169
+ if rolling is None and cross is not None
170
+ }
171
+ shared_plan_counts: dict[str, int] = {}
172
+ for plan in (*plans.values(), *cross_plans.values()):
173
+ if plan is not None:
174
+ shared_plan_counts[plan.node_id] = (
175
+ shared_plan_counts.get(plan.node_id, 0) + 1
176
+ )
177
+ nodes: list[dict[str, object]] = []
178
+ edges: list[dict[str, object]] = []
179
+ fanout_ids: dict[str, str] = {}
180
+ if fanout:
181
+ for value in program.inputs:
182
+ input_node = value._node
183
+ if input_node.digest not in consumed:
184
+ continue
185
+ input_name = _cstr(input_node.attr("name"))
186
+ schema = _schema_fields(input_node.attr("schema"))
187
+ pinned = (
188
+ input_node.digest in rolling_digests
189
+ or input_node.digest in direct_cross_section_digests
190
+ )
191
+ nodes.append(
192
+ _expression_node(
193
+ input_name,
194
+ [_quote_identifier(field.name) for field in schema],
195
+ None,
196
+ schema,
197
+ schema if pinned else None,
198
+ )
199
+ )
200
+ fanout_ids[input_node.digest] = input_name
201
+ for output_name, segment in segments:
202
+ rolling = plans[output_name]
203
+ cross = cross_plans[output_name]
204
+ env = dict(segment.env)
205
+ input_field_names = [
206
+ field.name for field in _schema_fields(segment.input_node.attr("schema"))
207
+ ]
208
+ upstream_id: str | None = None
209
+ final_predicate = segment.predicate
210
+ if (rolling is not None or cross is not None) and segment.predicate is not None:
211
+ state_plan = _required_segment_state_plan(rolling, cross)
212
+ prefilter_id = (
213
+ f"{state_plan.node_id}__prefilter"
214
+ if shared_plan_counts[state_plan.node_id] > 1
215
+ else f"{output_name}__cf_prefilter"
216
+ )
217
+ input_fields = _schema_fields(segment.input_node.attr("schema"))
218
+ nodes.append(
219
+ _expression_node(
220
+ prefilter_id,
221
+ [_quote_identifier(name) for name in input_field_names],
222
+ _sql(segment.predicate),
223
+ input_fields,
224
+ input_fields,
225
+ )
226
+ )
227
+ if fanout:
228
+ edges.append(
229
+ {
230
+ "source_node": fanout_ids[segment.input_node.digest],
231
+ "source_port": "output",
232
+ "target_node": prefilter_id,
233
+ "target_port": "input",
234
+ }
235
+ )
236
+ upstream_id = prefilter_id
237
+ if rolling is not None:
238
+ for stage in rolling.stages:
239
+ if stage.materialization_node is not None:
240
+ nodes.append(stage.materialization_node)
241
+ materialization_id = stage.materialization_node_id
242
+ if materialization_id is None:
243
+ raise RuntimeError("rolling materialization node has no id")
244
+ if upstream_id is not None:
245
+ edges.append(
246
+ {
247
+ "source_node": upstream_id,
248
+ "source_port": "output",
249
+ "target_node": materialization_id,
250
+ "target_port": "input",
251
+ }
252
+ )
253
+ elif fanout:
254
+ edges.append(
255
+ {
256
+ "source_node": fanout_ids[segment.input_node.digest],
257
+ "source_port": "output",
258
+ "target_node": materialization_id,
259
+ "target_port": "input",
260
+ }
261
+ )
262
+ upstream_id = materialization_id
263
+ nodes.append(stage.node)
264
+ if upstream_id is not None:
265
+ edges.append(
266
+ {
267
+ "source_node": upstream_id,
268
+ "source_port": "output",
269
+ "target_node": stage.node_id,
270
+ "target_port": "input",
271
+ }
272
+ )
273
+ elif fanout:
274
+ edges.append(
275
+ {
276
+ "source_node": fanout_ids[segment.input_node.digest],
277
+ "source_port": "output",
278
+ "target_node": stage.node_id,
279
+ "target_port": "input",
280
+ }
281
+ )
282
+ upstream_id = stage.node_id
283
+ env = dict(rolling.env)
284
+ input_field_names = list(rolling.input_field_names)
285
+ final_predicate = rolling.post_predicate
286
+ if cross is not None:
287
+ if cross.materialization_node is not None:
288
+ nodes.append(cross.materialization_node)
289
+ materialization_id = cross.materialization_node_id
290
+ if materialization_id is None:
291
+ raise RuntimeError("cross-section materialization node has no id")
292
+ if upstream_id is not None:
293
+ edges.append(
294
+ {
295
+ "source_node": upstream_id,
296
+ "source_port": "output",
297
+ "target_node": materialization_id,
298
+ "target_port": "input",
299
+ }
300
+ )
301
+ elif fanout:
302
+ edges.append(
303
+ {
304
+ "source_node": fanout_ids[segment.input_node.digest],
305
+ "source_port": "output",
306
+ "target_node": materialization_id,
307
+ "target_port": "input",
308
+ }
309
+ )
310
+ upstream_id = materialization_id
311
+ nodes.append(cross.node)
312
+ if upstream_id is not None:
313
+ edges.append(
314
+ {
315
+ "source_node": upstream_id,
316
+ "source_port": "output",
317
+ "target_node": cross.node_id,
318
+ "target_port": "input",
319
+ }
320
+ )
321
+ elif fanout:
322
+ edges.append(
323
+ {
324
+ "source_node": fanout_ids[segment.input_node.digest],
325
+ "source_port": "output",
326
+ "target_node": cross.node_id,
327
+ "target_port": "input",
328
+ }
329
+ )
330
+ upstream_id = cross.node_id
331
+ env = dict(cross.env)
332
+ input_field_names = list(cross.input_field_names)
333
+ final_predicate = cross.post_predicate
334
+ elif rolling is None and segment.post_predicate is not None:
335
+ final_predicate = (
336
+ segment.post_predicate
337
+ if segment.predicate is None
338
+ else build("and", (segment.predicate, segment.post_predicate), {})
339
+ )
340
+ reserved = frozenset(env)
341
+ fused = extract_common(
342
+ tuple((field, env[field]) for field in segment.fields),
343
+ final_predicate,
344
+ reserved,
345
+ )
346
+ cse_order = [name for tier in fused.tiers for name, _ in tier]
347
+ needed: set[str] = set()
348
+ for _, tree in fused.selects:
349
+ needed |= expression_refs(tree)
350
+ if fused.predicate is not None:
351
+ needed |= expression_refs(fused.predicate)
352
+ tier_items: list[list[str]] = []
353
+ for index in range(len(fused.tiers) - 1, -1, -1):
354
+ tier = fused.tiers[index]
355
+ defined = {name for name, _ in tier}
356
+ passthrough = needed - defined
357
+ items = [
358
+ _quote_identifier(field)
359
+ for field in input_field_names
360
+ if field in passthrough
361
+ ]
362
+ items += [
363
+ _quote_identifier(name) for name in cse_order if name in passthrough
364
+ ]
365
+ items += [
366
+ f"{_sql(tree)} AS {_quote_identifier(name)}" for name, tree in tier
367
+ ]
368
+ tier_items.append(items)
369
+ needed = set(passthrough)
370
+ for _, tree in tier:
371
+ needed |= expression_refs(tree)
372
+ tier_items.reverse()
373
+ stage_ids = [
374
+ f"{output_name}__cf_cse_{index}" for index in range(1, len(fused.tiers) + 1)
375
+ ] + [output_name]
376
+ stage_selects = [
377
+ *tier_items,
378
+ [_select_item(field, tree) for field, tree in fused.selects],
379
+ ]
380
+ for position, (node_id, select) in enumerate(
381
+ zip(stage_ids, stage_selects, strict=True)
382
+ ):
383
+ filter_sql = None
384
+ if position == len(stage_ids) - 1 and fused.predicate is not None:
385
+ filter_sql = _sql(fused.predicate)
386
+ if position == 0:
387
+ if upstream_id is not None:
388
+ edges.append(
389
+ {
390
+ "source_node": upstream_id,
391
+ "source_port": "output",
392
+ "target_node": node_id,
393
+ "target_port": "input",
394
+ }
395
+ )
396
+ input_schema = (
397
+ cross.output_fields
398
+ if cross is not None
399
+ else rolling.output_fields
400
+ if rolling is not None
401
+ else None
402
+ )
403
+ elif fanout:
404
+ edges.append(
405
+ {
406
+ "source_node": fanout_ids[segment.input_node.digest],
407
+ "source_port": "output",
408
+ "target_node": node_id,
409
+ "target_port": "input",
410
+ }
411
+ )
412
+ input_schema = None
413
+ else:
414
+ input_schema = _schema_fields(segment.input_node.attr("schema"))
415
+ else:
416
+ edges.append(
417
+ {
418
+ "source_node": stage_ids[position - 1],
419
+ "source_port": "output",
420
+ "target_node": node_id,
421
+ "target_port": "input",
422
+ }
423
+ )
424
+ input_schema = None
425
+ nodes.append(_expression_node(node_id, select, filter_sql, input_schema))
426
+ nodes, edges = _deduplicate_node_ids(nodes, edges)
427
+ nodes, edges = _deduplicate_pure_expression_nodes(
428
+ nodes,
429
+ edges,
430
+ frozenset(output_name for output_name, _ in program.outputs),
431
+ )
432
+ return _project_document(program.name, mode, nodes, edges)
433
+
434
+
435
+ def _require_runtime(runtime: object, entry: str, /) -> Runtime:
436
+ if not isinstance(runtime, Runtime):
437
+ raise TypeError(
438
+ f"{entry} requires an explicit calc_flow Runtime; got {type_name(runtime)}"
439
+ )
440
+ return runtime
441
+
442
+
443
+ def _check_expression_capability(
444
+ program: Program,
445
+ runtime: Runtime,
446
+ mode: str,
447
+ /,
448
+ ) -> tuple[_Analyzer, RuntimeCapabilities]:
449
+ analyzer, capabilities = _run(program, runtime, mode)
450
+ issues = analyzer.issues
451
+ if issues:
452
+ first = issues[0]
453
+ errors.raise_compile(first.path, first.code, first.message)
454
+ for operator in capabilities.operators:
455
+ if operator.kind != "expression":
456
+ continue
457
+ _require_mode_support(program, operator, mode)
458
+ _require_stream_facts(program, operator, mode)
459
+ return analyzer, capabilities
460
+ errors.raise_compile(
461
+ program.name,
462
+ errors.CAPABILITY_MISMATCH,
463
+ "the capability snapshot does not offer the expression operator",
464
+ )
465
+
466
+
467
+ def _require_mode_support(program: Program, operator: object, mode: str, /) -> None:
468
+ if mode not in operator.modes: # type: ignore[attr-defined]
469
+ errors.raise_compile(
470
+ program.name,
471
+ errors.CAPABILITY_MISMATCH,
472
+ f"the expression operator does not support {mode} mode in the"
473
+ " selected capability snapshot",
474
+ )
475
+
476
+
477
+ def _require_stream_facts(program: Program, operator: object, mode: str, /) -> None:
478
+ if mode == "stream" and (
479
+ operator.finality == "unproven" # type: ignore[attr-defined]
480
+ or not operator.microbatch_invariant # type: ignore[attr-defined]
481
+ or not operator.deterministic # type: ignore[attr-defined]
482
+ or not operator.replay_safe # type: ignore[attr-defined]
483
+ ):
484
+ errors.raise_compile(
485
+ program.name,
486
+ errors.CAPABILITY_MISMATCH,
487
+ "the expression operator does not prove stream lifecycle facts"
488
+ " in the selected capability snapshot",
489
+ )
490
+
491
+
492
+ def _program_needs_rolling(program: Program, /) -> bool:
493
+ return any(True for _, value in program.outputs for _ in _find_rolling(value._node))
494
+
495
+
496
+ def _program_needs_cross_section(program: Program, /) -> bool:
497
+ return any(
498
+ True for _, value in program.outputs for _ in _find_cross_section(value._node)
499
+ )
500
+
501
+
502
+ def _program_needs_stream_join(program: Program, /) -> bool:
503
+ return bool(_stream_join_nodes(program))
504
+
505
+
506
+ def _stream_join_ports(operator: object, /) -> bool:
507
+ inputs = tuple(
508
+ (port.name, port.kind, port.required) for port in operator.input_ports
509
+ )
510
+ outputs = tuple(
511
+ (port.name, port.kind, port.required) for port in operator.output_ports
512
+ )
513
+ return inputs == (("left", "table", True), ("right", "table", True)) and (
514
+ outputs == (("output", "table", True),)
515
+ )
516
+
517
+
518
+ def _positive_state_version(value: object, /) -> bool:
519
+ return isinstance(value, int) and value > 0
520
+
521
+
522
+ def _stream_join_capability_facts(operator: object, mode: str, /) -> tuple[bool, ...]:
523
+ return (
524
+ operator.version == "1",
525
+ mode in operator.modes,
526
+ _stream_join_ports(operator),
527
+ operator.requires_watermark,
528
+ operator.stateful,
529
+ operator.checkpoint_support == "checkpointed_stateful",
530
+ _positive_state_version(operator.state_version),
531
+ operator.deterministic,
532
+ operator.replay_safe,
533
+ )
534
+
535
+
536
+ def _check_stream_join_capability(
537
+ program: Program,
538
+ capabilities: RuntimeCapabilities,
539
+ mode: str,
540
+ /,
541
+ ) -> None:
542
+ for operator in capabilities.operators:
543
+ if operator.kind != "stream_join":
544
+ continue
545
+ if not all(_stream_join_capability_facts(operator, mode)):
546
+ errors.raise_compile(
547
+ program.name,
548
+ errors.CAPABILITY_MISMATCH,
549
+ "the stream_join operator does not prove the required stream"
550
+ " ports, watermark, checkpoint, determinism, and replay facts",
551
+ )
552
+ return
553
+ errors.raise_compile(
554
+ program.name,
555
+ errors.CAPABILITY_MISMATCH,
556
+ "the capability snapshot does not offer stream_join@1",
557
+ )
558
+
559
+
560
+ # The cross-section gate mirrors the rolling gate over the frozen
561
+ # group-final stream lifecycle facts.
562
+ def _cross_section_checkpoint_capability(operator: object, /) -> bool:
563
+ return (
564
+ operator.stateful
565
+ and operator.checkpoint_support == "checkpointed_stateful"
566
+ and isinstance(operator.state_version, int)
567
+ and operator.state_version > 0
568
+ )
569
+
570
+
571
+ def _cross_section_stream_capability(operator: object, /) -> bool:
572
+ return (
573
+ operator.finality != "unproven"
574
+ and operator.microbatch_invariant
575
+ and _cross_section_checkpoint_capability(operator)
576
+ and operator.deterministic
577
+ and operator.replay_safe
578
+ )
579
+
580
+
581
+ def _check_cross_section_capability(
582
+ program: Program,
583
+ capabilities: object,
584
+ mode: str,
585
+ /,
586
+ ) -> None:
587
+ for operator in capabilities.operators:
588
+ if operator.kind != "cross_section":
589
+ continue
590
+ if mode not in operator.modes:
591
+ errors.raise_compile(
592
+ program.name,
593
+ errors.CAPABILITY_MISMATCH,
594
+ f"the cross-section operator does not support {mode} mode in"
595
+ " the selected capability snapshot",
596
+ )
597
+ if mode == "stream" and not _cross_section_stream_capability(operator):
598
+ errors.raise_compile(
599
+ program.name,
600
+ errors.CAPABILITY_MISMATCH,
601
+ "the cross-section operator does not prove stream lifecycle"
602
+ " facts in the selected capability snapshot",
603
+ )
604
+ return
605
+ errors.raise_compile(
606
+ program.name,
607
+ errors.CAPABILITY_MISMATCH,
608
+ "the capability snapshot does not offer the cross-section operator",
609
+ )
610
+
611
+
612
+ # The capability gate conjoins the frozen stream lifecycle facts; every
613
+ # fact fails with the same stable capability_mismatch code.
614
+ def _check_rolling_capability(
615
+ program: Program,
616
+ capabilities: object,
617
+ mode: str,
618
+ /,
619
+ ) -> None:
620
+ # #lizard forgives
621
+ for operator in capabilities.operators:
622
+ if operator.kind != "rolling":
623
+ continue
624
+ if mode not in operator.modes:
625
+ errors.raise_compile(
626
+ program.name,
627
+ errors.CAPABILITY_MISMATCH,
628
+ f"the rolling operator does not support {mode} mode in the"
629
+ " selected capability snapshot",
630
+ )
631
+ if mode == "stream" and (
632
+ operator.finality == "unproven"
633
+ or not operator.stateful
634
+ or not operator.microbatch_invariant
635
+ or operator.checkpoint_support != "checkpointed_stateful"
636
+ or not isinstance(operator.state_version, int)
637
+ or operator.state_version <= 0
638
+ or not operator.deterministic
639
+ or not operator.replay_safe
640
+ ):
641
+ errors.raise_compile(
642
+ program.name,
643
+ errors.CAPABILITY_MISMATCH,
644
+ "the rolling operator does not prove stream lifecycle facts"
645
+ " in the selected capability snapshot",
646
+ )
647
+ return
648
+ errors.raise_compile(
649
+ program.name,
650
+ errors.CAPABILITY_MISMATCH,
651
+ "the capability snapshot does not offer the rolling operator",
652
+ )
653
+
654
+
655
+ def lower_program_document(
656
+ program: Program,
657
+ runtime: Runtime,
658
+ mode: str,
659
+ /,
660
+ *,
661
+ allowed_lateness_micros: int = 0,
662
+ late_policy: str = "error",
663
+ ) -> dict[str, object]:
664
+ """Analyze and lower one program to its strict project-v3 document.
665
+
666
+ The lateness arguments are validated whenever the program contains
667
+ rolling or cross-section primitives; row-local programs do not consume
668
+ them.
669
+ """
670
+
671
+ selected = _require_runtime(runtime, "lower_program_document")
672
+ mode_value = _require_mode(mode)
673
+ analyzer, capabilities = _check_expression_capability(program, selected, mode_value)
674
+ if _program_needs_stream_join(program):
675
+ _check_stream_join_capability(program, capabilities, mode_value)
676
+ if _program_needs_rolling(program) or _program_needs_cross_section(program):
677
+ _validate_lateness(allowed_lateness_micros, late_policy)
678
+ if _program_needs_rolling(program):
679
+ _check_rolling_capability(program, capabilities, mode_value)
680
+ if _program_needs_cross_section(program):
681
+ _check_cross_section_capability(program, capabilities, mode_value)
682
+ join_project = _lower_stream_join_program(
683
+ program,
684
+ analyzer,
685
+ mode_value,
686
+ allowed_lateness_micros,
687
+ late_policy,
688
+ )
689
+ if join_project is not None:
690
+ return join_project
691
+ return _lower_program(program, mode_value, allowed_lateness_micros, late_policy)
692
+
693
+
694
+ def _cache_graph_nodes(document: dict[str, object], /) -> list[dict[str, object]]:
695
+ graph = document["graph"]
696
+ return graph["nodes"] # type: ignore[index,return-value]
697
+
698
+
699
+ def _cache_operator_versions(
700
+ nodes: list[dict[str, object]], capabilities: RuntimeCapabilities, /
701
+ ) -> tuple[tuple[str, str], ...]:
702
+ operator_kinds = {
703
+ node["operator"]["kind"] # type: ignore[index]
704
+ for node in nodes
705
+ if node["operator"]["kind"] != "external" # type: ignore[index]
706
+ }
707
+ return tuple(
708
+ sorted(
709
+ (operator.kind, operator.version)
710
+ for operator in capabilities.operators
711
+ if operator.kind in operator_kinds
712
+ )
713
+ )
714
+
715
+
716
+ def _cache_provider_versions(
717
+ nodes: list[dict[str, object]], /
718
+ ) -> tuple[tuple[str, str, str], ...]:
719
+ versions: set[tuple[str, str, str]] = set()
720
+ for node in nodes:
721
+ operator = node["operator"]
722
+ if operator["kind"] == "external": # type: ignore[index]
723
+ versions.add( # type: ignore[arg-type]
724
+ (operator["provider"], operator["name"], operator["version"]) # type: ignore[index]
725
+ )
726
+ return tuple(sorted(versions))
727
+
728
+
729
+ def _cache_udf_versions(
730
+ nodes: list[dict[str, object]], /
731
+ ) -> tuple[tuple[str, str, str], ...]:
732
+ versions: set[tuple[str, str, str]] = set()
733
+ for node in nodes:
734
+ operator = node["operator"]
735
+ for udf in operator.get("udfs", ()): # type: ignore[union-attr]
736
+ if isinstance(udf, dict):
737
+ versions.add((udf["provider"], udf["name"], udf["version"])) # type: ignore[arg-type]
738
+ return tuple(sorted(versions))
739
+
740
+
741
+ def _compile_cache_key(
742
+ program: Program,
743
+ mode: str,
744
+ document: dict[str, object],
745
+ capabilities: RuntimeCapabilities,
746
+ allowed_lateness_micros: int,
747
+ late_policy: str,
748
+ /,
749
+ ) -> _CompileCacheKey:
750
+ nodes = _cache_graph_nodes(document)
751
+ return _CompileCacheKey(
752
+ program_fingerprint=program.fingerprint,
753
+ mode=mode,
754
+ input_declarations=tuple(
755
+ value._node.node_bytes.hex() for value in program.inputs
756
+ ),
757
+ capability_schema_version=capabilities.schema_version,
758
+ capability_session_id=capabilities.scope.session_id,
759
+ capability_revision=capabilities.scope.revision,
760
+ operator_versions=_cache_operator_versions(nodes, capabilities),
761
+ provider_versions=_cache_provider_versions(nodes),
762
+ udf_versions=_cache_udf_versions(nodes),
763
+ allowed_lateness_micros=allowed_lateness_micros,
764
+ late_policy=late_policy,
765
+ )
766
+
767
+
768
+ def compile_program_batch(program: Program, runtime: object, /) -> BatchExecutionPlan:
769
+ """Lower one program to a strict project-v3 batch plan."""
770
+
771
+ selected = _require_runtime(runtime, "compile_batch")
772
+ document = lower_program_document(program, selected, "batch")
773
+ capabilities = selected.capabilities()
774
+ key = _compile_cache_key(program, "batch", document, capabilities, 0, "error")
775
+ return selected._cached_symbolic_compile(
776
+ key, lambda: selected.compile_batch_project(_canonical(document))
777
+ ) # type: ignore[return-value]
778
+
779
+
780
+ def compile_program_stream(
781
+ program: Program,
782
+ runtime: object,
783
+ allowed_lateness_micros: object,
784
+ late_policy: object,
785
+ /,
786
+ ) -> StreamExecutionPlan:
787
+ """Lower one program to a strict project-v3 continuous plan.
788
+
789
+ The validated lateness arguments are written into every lowered rolling
790
+ node; row-local programs are unaffected by them.
791
+ """
792
+
793
+ selected = _require_runtime(runtime, "compile_stream")
794
+ _validate_lateness(allowed_lateness_micros, late_policy)
795
+ document = lower_program_document(
796
+ program,
797
+ selected,
798
+ "stream",
799
+ allowed_lateness_micros=allowed_lateness_micros,
800
+ late_policy=late_policy,
801
+ )
802
+ capabilities = selected.capabilities()
803
+ key = _compile_cache_key(
804
+ program,
805
+ "stream",
806
+ document,
807
+ capabilities,
808
+ allowed_lateness_micros,
809
+ late_policy,
810
+ )
811
+ return selected._cached_symbolic_compile(
812
+ key,
813
+ lambda: selected._compile_stream_graph_project(
814
+ _canonical(document), requirements=StreamRequirements()
815
+ ),
816
+ ) # type: ignore[return-value]
817
+
818
+
819
+ def _validate_lateness(allowed_lateness_micros: object, late_policy: object, /) -> None:
820
+ if type(allowed_lateness_micros) is not int:
821
+ raise TypeError(
822
+ "allowed_lateness_micros must be an exact int; got"
823
+ f" {type_name(allowed_lateness_micros)}"
824
+ )
825
+ if allowed_lateness_micros < 0:
826
+ raise ValueError(
827
+ "allowed_lateness_micros: invalid_literal: must be non-negative"
828
+ )
829
+ if allowed_lateness_micros > _U64_MAX:
830
+ raise ValueError(
831
+ "allowed_lateness_micros: invalid_literal: must fit the unsigned"
832
+ " 64-bit microsecond range"
833
+ )
834
+ if type(late_policy) is not str:
835
+ raise TypeError(f"late_policy must be a string; got {type_name(late_policy)}")
836
+ if late_policy not in ("error", "drop"):
837
+ raise ValueError(
838
+ "late_policy: invalid_literal: must be 'error' or"
839
+ f" 'drop'; got {late_policy!r}"
840
+ )