refactor: derive context from foreach control regions

This commit is contained in:
lda
2026-09-04 07:18:03 +07:00 Verified
parent 6a5d886962
commit 4bbd9f9650
2 changed files with 95 additions and 108 deletions
+59 -82
View File
@@ -1,11 +1,14 @@
from __future__ import annotations from __future__ import annotations
from collections import deque
from collections.abc import Mapping from collections.abc import Mapping
from copy import deepcopy from copy import deepcopy
from dataclasses import dataclass from dataclasses import dataclass
from typing import Literal from typing import Literal
from wf_core.analysis.control_regions import (
ForeachOwnerStack,
analyze_control_regions,
)
from wf_core.context_contracts import ( from wf_core.context_contracts import (
STANDARD_CONTEXT_FIELDS, STANDARD_CONTEXT_FIELDS,
ContextFieldContract, ContextFieldContract,
@@ -80,6 +83,13 @@ def context_analysis_warnings(workflow: Workflow) -> tuple[str, ...]:
def _analyze(workflow: Workflow) -> _ContextAnalysis: def _analyze(workflow: Workflow) -> _ContextAnalysis:
"""Derive context contracts from static foreach control regions.
Each unambiguous node use has exactly one owner stack; its active foreach
is the final stack item. A canonical return edge pops the stack, so the
controller itself stays in the outer context. Conflicted nodes receive no
foreach fields.
"""
nodes = {node.id: node for node in workflow.nodes} nodes = {node.id: node for node in workflow.nodes}
foreach_nodes = { foreach_nodes = {
node.id: node for node in workflow.nodes if isinstance(node, ForeachNode) node.id: node for node in workflow.nodes if isinstance(node, ForeachNode)
@@ -107,35 +117,20 @@ def _analyze(workflow: Workflow) -> _ContextAnalysis:
warnings.add(f"workflow start targets missing node {workflow.start!r}") warnings.add(f"workflow start targets missing node {workflow.start!r}")
return _ContextAnalysis({}, tuple(warnings.values)) return _ContextAnalysis({}, tuple(warnings.values))
scopes_by_node: dict[str, set[FrameScope]] = {} analysis = analyze_control_regions(workflow)
pending: deque[tuple[str, FrameScope]] = deque([(workflow.start, None)]) for issue in analysis.issues:
visited: set[tuple[str, FrameScope]] = set() warnings.add(
while pending: f"control region {issue.kind.value} at {issue.path}: {issue.message}"
node_id, active_scope = pending.popleft() )
state = (node_id, active_scope)
if state in visited:
continue
visited.add(state)
node = nodes.get(node_id)
if node is None:
continue
scopes_by_node.setdefault(node_id, set()).add(active_scope)
for edge in edges_by_node.get(node_id, []):
if edge.to == END or edge.to not in nodes:
continue
next_scope = active_scope
if isinstance(node, ForeachNode) and edge.outcome == "loop":
next_scope = node.id
pending.append((edge.to, next_scope))
fields_by_node: dict[str, tuple[ContextFieldAvailability, ...]] = {} fields_by_node: dict[str, tuple[ContextFieldAvailability, ...]] = {}
for node_id, scopes in scopes_by_node.items(): for node_id, stack in analysis.owner_stack_by_node.items():
active_foreach_id = stack[-1] if stack else None
fields_by_node[node_id] = _available_fields( fields_by_node[node_id] = _available_fields(
workflow, workflow,
foreach_nodes, foreach_nodes,
scopes, analysis.owner_stack_by_node,
scopes_by_node, active_foreach_id,
) )
return _ContextAnalysis(fields_by_node, tuple(warnings.values)) return _ContextAnalysis(fields_by_node, tuple(warnings.values))
@@ -143,83 +138,66 @@ def _analyze(workflow: Workflow) -> _ContextAnalysis:
def _available_fields( def _available_fields(
workflow: Workflow, workflow: Workflow,
foreach_nodes: Mapping[str, ForeachNode], foreach_nodes: Mapping[str, ForeachNode],
scopes: set[FrameScope], owner_stack_by_node: Mapping[str, ForeachOwnerStack],
scopes_by_node: Mapping[str, set[FrameScope]], active_scope: FrameScope,
) -> tuple[ContextFieldAvailability, ...]: ) -> tuple[ContextFieldAvailability, ...]:
fields_by_name: dict[str, ContextFieldContract] = {} """Return contracts for one static owner stack; all are guaranteed.
scopes_by_field: dict[str, set[FrameScope]] = {}
for scope in sorted(scopes, key=lambda value: value or ""): A single node use has one control region, so foreach fields are either
contracts = STANDARD_CONTEXT_FIELDS present (inside a body) or absent (outside). Conditional availability is
if scope is not None: not used to represent multiple owner stacks.
foreach = foreach_nodes.get(scope) """
contracts = list(STANDARD_CONTEXT_FIELDS)
if active_scope is not None:
foreach = foreach_nodes.get(active_scope)
if foreach is not None: if foreach is not None:
contracts = ( contracts.extend(
*contracts, foreach_context_fields(
*foreach_context_fields(
foreach.as_, foreach.as_,
_foreach_item_schema( _foreach_item_schema(
workflow, workflow,
foreach, foreach,
scopes_by_node.get(foreach.id, {None}),
foreach_nodes, foreach_nodes,
scopes_by_node, owner_stack_by_node,
),
), ),
) )
for contract in contracts: )
fields_by_name.setdefault( return tuple(
contract.name, ContextFieldAvailability(
ContextFieldContract( contract=ContextFieldContract(
contract.name, contract.name,
deepcopy(contract.schema), deepcopy(contract.schema),
contract.description, contract.description,
), ),
availability="available",
) )
scopes_by_field.setdefault(contract.name, set()).add(scope) for contract in contracts
field_count = len(scopes)
result: list[ContextFieldAvailability] = []
for contract in fields_by_name.values():
field_scopes = scopes_by_field[contract.name]
availability: ContextAvailability = (
"available" if len(field_scopes) == field_count else "conditional"
) )
reason = None
if availability == "conditional":
reason = "Available only in some reachable execution frames."
result.append(
ContextFieldAvailability(
contract=contract,
availability=availability,
reason=reason,
)
)
return tuple(result)
def _foreach_item_schema( def _foreach_item_schema(
workflow: Workflow, workflow: Workflow,
foreach: ForeachNode, foreach: ForeachNode,
source_scopes: set[FrameScope],
foreach_nodes: Mapping[str, ForeachNode], foreach_nodes: Mapping[str, ForeachNode],
scopes_by_node: Mapping[str, set[FrameScope]], owner_stack_by_node: Mapping[str, ForeachOwnerStack],
) -> ContextSchema: ) -> ContextSchema:
source_schemas = [ """Resolve one controller's item schema in its own static context.
_schema_at_path(
An inner foreach may declare ``over="context.outer_item"``; the lookup
uses the controller's own owner stack, not the inner body stack.
"""
controller_stack = owner_stack_by_node.get(foreach.id)
if controller_stack is None:
return {}
controller_scope: FrameScope = controller_stack[-1] if controller_stack else None
source_schema = _schema_at_path(
workflow, workflow,
foreach.over.root, foreach.over.root,
foreach.over.parts, foreach.over.parts,
source_scope, controller_scope,
foreach_nodes, foreach_nodes,
scopes_by_node, owner_stack_by_node,
) )
for source_scope in sorted(source_scopes, key=lambda value: value or "")
]
if not source_schemas or any(
schema != source_schemas[0] for schema in source_schemas
):
return {}
source_schema = source_schemas[0]
if not isinstance(source_schema, Mapping): if not isinstance(source_schema, Mapping):
return {} return {}
source_type = source_schema.get("type") source_type = source_schema.get("type")
@@ -235,7 +213,7 @@ def _foreach_item_schema(
workflow, workflow,
foreach.over.root, foreach.over.root,
foreach_nodes=foreach_nodes, foreach_nodes=foreach_nodes,
scopes_by_node=scopes_by_node, owner_stack_by_node=owner_stack_by_node,
), ),
items, items,
) )
@@ -250,7 +228,7 @@ def _schema_at_path(
parts: tuple[str, ...], parts: tuple[str, ...],
active_scope: FrameScope, active_scope: FrameScope,
foreach_nodes: Mapping[str, ForeachNode], foreach_nodes: Mapping[str, ForeachNode],
scopes_by_node: Mapping[str, set[FrameScope]], owner_stack_by_node: Mapping[str, ForeachOwnerStack],
) -> Mapping[str, object] | None: ) -> Mapping[str, object] | None:
try: try:
schema_document = _schema_document( schema_document = _schema_document(
@@ -258,7 +236,7 @@ def _schema_at_path(
root, root,
active_scope=active_scope, active_scope=active_scope,
foreach_nodes=foreach_nodes, foreach_nodes=foreach_nodes,
scopes_by_node=scopes_by_node, owner_stack_by_node=owner_stack_by_node,
) )
current: object = schema_document current: object = schema_document
for part in parts: for part in parts:
@@ -282,7 +260,7 @@ def _schema_document(
*, *,
active_scope: FrameScope = None, active_scope: FrameScope = None,
foreach_nodes: Mapping[str, ForeachNode] | None = None, foreach_nodes: Mapping[str, ForeachNode] | None = None,
scopes_by_node: Mapping[str, set[FrameScope]] | None = None, owner_stack_by_node: Mapping[str, ForeachOwnerStack] | None = None,
) -> Mapping[str, object]: ) -> Mapping[str, object]:
if root == "input": if root == "input":
return workflow.input_schema.model_dump(mode="json", exclude_none=True) return workflow.input_schema.model_dump(mode="json", exclude_none=True)
@@ -295,7 +273,7 @@ def _schema_document(
if ( if (
active_scope is not None active_scope is not None
and foreach_nodes is not None and foreach_nodes is not None
and scopes_by_node is not None and owner_stack_by_node is not None
): ):
foreach = foreach_nodes.get(active_scope) foreach = foreach_nodes.get(active_scope)
if foreach is not None: if foreach is not None:
@@ -307,9 +285,8 @@ def _schema_document(
_foreach_item_schema( _foreach_item_schema(
workflow, workflow,
foreach, foreach,
scopes_by_node.get(foreach.id, {None}),
foreach_nodes, foreach_nodes,
scopes_by_node, owner_stack_by_node,
), ),
) )
} }
+22 -12
View File
@@ -146,7 +146,7 @@ def test_serial_and_concurrent_foreach_expose_the_same_scoped_context() -> None:
edges=[ edges=[
{"from": "each", "outcome": "loop", "to": "body"}, {"from": "each", "outcome": "loop", "to": "body"},
{"from": "each", "outcome": "done", "to": "tail"}, {"from": "each", "outcome": "done", "to": "tail"},
{"from": "body", "outcome": "ok", "to": END}, {"from": "body", "outcome": "ok", "to": "each"},
{"from": "tail", "outcome": "ok", "to": END}, {"from": "tail", "outcome": "ok", "to": END},
], ],
) )
@@ -164,7 +164,7 @@ def test_foreach_item_schema_and_configured_alias_are_reported() -> None:
nodes=[_foreach("each", alias="record"), _node("body")], nodes=[_foreach("each", alias="record"), _node("body")],
edges=[ edges=[
{"from": "each", "outcome": "loop", "to": "body"}, {"from": "each", "outcome": "loop", "to": "body"},
{"from": "body", "outcome": "ok", "to": END}, {"from": "body", "outcome": "ok", "to": "each"},
{"from": "each", "outcome": "done", "to": END}, {"from": "each", "outcome": "done", "to": END},
], ],
) )
@@ -181,7 +181,7 @@ def test_foreach_item_schema_resolves_bounded_local_array_reference() -> None:
nodes=[_foreach("each", alias="record"), _node("body")], nodes=[_foreach("each", alias="record"), _node("body")],
edges=[ edges=[
{"from": "each", "outcome": "loop", "to": "body"}, {"from": "each", "outcome": "loop", "to": "body"},
{"from": "body", "outcome": "ok", "to": END}, {"from": "body", "outcome": "ok", "to": "each"},
{"from": "each", "outcome": "done", "to": END}, {"from": "each", "outcome": "done", "to": END},
], ],
state_schema={ state_schema={
@@ -209,7 +209,7 @@ def test_foreach_item_schema_resolves_bounded_local_array_reference() -> None:
assert fields["record"].contract.schema["properties"] == {"id": {"type": "string"}} assert fields["record"].contract.schema["properties"] == {"id": {"type": "string"}}
def test_only_foreach_reachable_node_has_available_context() -> None: def test_region_conflicted_node_receives_no_guaranteed_foreach_fields() -> None:
workflow = _workflow( workflow = _workflow(
start="start", start="start",
nodes=[_node("start"), _foreach("each", alias="item"), _node("body")], nodes=[_node("start"), _foreach("each", alias="item"), _node("body")],
@@ -218,12 +218,18 @@ def test_only_foreach_reachable_node_has_available_context() -> None:
{"from": "start", "outcome": "loop", "to": "each"}, {"from": "start", "outcome": "loop", "to": "each"},
{"from": "each", "outcome": "loop", "to": "body"}, {"from": "each", "outcome": "loop", "to": "body"},
{"from": "each", "outcome": "done", "to": END}, {"from": "each", "outcome": "done", "to": END},
{"from": "body", "outcome": "ok", "to": END}, {"from": "body", "outcome": "ok", "to": "each"},
], ],
) )
assert _field_map(workflow, "body")["item"].availability == "conditional" fields = context_fields_by_node(workflow)
assert _field_map(workflow, "body")["item"].reason assert "body" not in fields or "item" not in {
field.contract.name for field in fields.get("body", ())
}
warnings = context_analysis_warnings(workflow)
assert any(
"control region" in warning or "conflict" in warning for warning in warnings
)
def test_nested_foreach_replaces_inner_scope_and_restores_outer_scope() -> None: def test_nested_foreach_replaces_inner_scope_and_restores_outer_scope() -> None:
@@ -239,8 +245,8 @@ def test_nested_foreach_replaces_inner_scope_and_restores_outer_scope() -> None:
{"from": "outer", "outcome": "loop", "to": "inner"}, {"from": "outer", "outcome": "loop", "to": "inner"},
{"from": "inner", "outcome": "loop", "to": "inner_body"}, {"from": "inner", "outcome": "loop", "to": "inner_body"},
{"from": "inner", "outcome": "done", "to": "after_inner"}, {"from": "inner", "outcome": "done", "to": "after_inner"},
{"from": "inner_body", "outcome": "ok", "to": END}, {"from": "inner_body", "outcome": "ok", "to": "inner"},
{"from": "after_inner", "outcome": "ok", "to": END}, {"from": "after_inner", "outcome": "ok", "to": "outer"},
{"from": "outer", "outcome": "done", "to": END}, {"from": "outer", "outcome": "done", "to": END},
], ],
state_schema={ state_schema={
@@ -267,12 +273,14 @@ def test_nested_foreach_preserves_context_backed_item_schema() -> None:
_foreach("outer", alias="outer_item"), _foreach("outer", alias="outer_item"),
_foreach("inner", alias="inner_item", over="context.outer_item"), _foreach("inner", alias="inner_item", over="context.outer_item"),
_node("inner_body"), _node("inner_body"),
_node("after_inner"),
], ],
edges=[ edges=[
{"from": "outer", "outcome": "loop", "to": "inner"}, {"from": "outer", "outcome": "loop", "to": "inner"},
{"from": "inner", "outcome": "loop", "to": "inner_body"}, {"from": "inner", "outcome": "loop", "to": "inner_body"},
{"from": "inner", "outcome": "done", "to": END}, {"from": "inner", "outcome": "done", "to": "after_inner"},
{"from": "inner_body", "outcome": "ok", "to": END}, {"from": "inner_body", "outcome": "ok", "to": "inner"},
{"from": "after_inner", "outcome": "ok", "to": "outer"},
{"from": "outer", "outcome": "done", "to": END}, {"from": "outer", "outcome": "done", "to": END},
], ],
state_schema={ state_schema={
@@ -350,4 +358,6 @@ def test_scoped_cycle_terminates_and_preserves_scoped_field_availability() -> No
fields = context_fields_by_node(workflow) fields = context_fields_by_node(workflow)
assert fields["body"] assert fields["body"]
assert _field_map(workflow, "body")["item"].availability == "available" assert _field_map(workflow, "body")["item"].availability == "available"
assert _field_map(workflow, "each")["item"].availability == "conditional" # A canonical back-edge pops the item stack, so the controller itself
# stays in the outer region and exposes no item alias.
assert "item" not in _field_map(workflow, "each")