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 collections import deque
from collections.abc import Mapping
from copy import deepcopy
from dataclasses import dataclass
from typing import Literal
from wf_core.analysis.control_regions import (
ForeachOwnerStack,
analyze_control_regions,
)
from wf_core.context_contracts import (
STANDARD_CONTEXT_FIELDS,
ContextFieldContract,
@@ -80,6 +83,13 @@ def context_analysis_warnings(workflow: Workflow) -> tuple[str, ...]:
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}
foreach_nodes = {
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}")
return _ContextAnalysis({}, tuple(warnings.values))
scopes_by_node: dict[str, set[FrameScope]] = {}
pending: deque[tuple[str, FrameScope]] = deque([(workflow.start, None)])
visited: set[tuple[str, FrameScope]] = set()
while pending:
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))
analysis = analyze_control_regions(workflow)
for issue in analysis.issues:
warnings.add(
f"control region {issue.kind.value} at {issue.path}: {issue.message}"
)
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(
workflow,
foreach_nodes,
scopes,
scopes_by_node,
analysis.owner_stack_by_node,
active_foreach_id,
)
return _ContextAnalysis(fields_by_node, tuple(warnings.values))
@@ -143,83 +138,66 @@ def _analyze(workflow: Workflow) -> _ContextAnalysis:
def _available_fields(
workflow: Workflow,
foreach_nodes: Mapping[str, ForeachNode],
scopes: set[FrameScope],
scopes_by_node: Mapping[str, set[FrameScope]],
owner_stack_by_node: Mapping[str, ForeachOwnerStack],
active_scope: FrameScope,
) -> tuple[ContextFieldAvailability, ...]:
fields_by_name: dict[str, ContextFieldContract] = {}
scopes_by_field: dict[str, set[FrameScope]] = {}
for scope in sorted(scopes, key=lambda value: value or ""):
contracts = STANDARD_CONTEXT_FIELDS
if scope is not None:
foreach = foreach_nodes.get(scope)
"""Return contracts for one static owner stack; all are guaranteed.
A single node use has one control region, so foreach fields are either
present (inside a body) or absent (outside). Conditional availability is
not used to represent multiple owner stacks.
"""
contracts = list(STANDARD_CONTEXT_FIELDS)
if active_scope is not None:
foreach = foreach_nodes.get(active_scope)
if foreach is not None:
contracts = (
*contracts,
*foreach_context_fields(
contracts.extend(
foreach_context_fields(
foreach.as_,
_foreach_item_schema(
workflow,
foreach,
scopes_by_node.get(foreach.id, {None}),
foreach_nodes,
scopes_by_node,
),
owner_stack_by_node,
),
)
for contract in contracts:
fields_by_name.setdefault(
contract.name,
ContextFieldContract(
)
return tuple(
ContextFieldAvailability(
contract=ContextFieldContract(
contract.name,
deepcopy(contract.schema),
contract.description,
),
availability="available",
)
scopes_by_field.setdefault(contract.name, set()).add(scope)
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"
for contract in contracts
)
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(
workflow: Workflow,
foreach: ForeachNode,
source_scopes: set[FrameScope],
foreach_nodes: Mapping[str, ForeachNode],
scopes_by_node: Mapping[str, set[FrameScope]],
owner_stack_by_node: Mapping[str, ForeachOwnerStack],
) -> ContextSchema:
source_schemas = [
_schema_at_path(
"""Resolve one controller's item schema in its own static context.
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,
foreach.over.root,
foreach.over.parts,
source_scope,
controller_scope,
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):
return {}
source_type = source_schema.get("type")
@@ -235,7 +213,7 @@ def _foreach_item_schema(
workflow,
foreach.over.root,
foreach_nodes=foreach_nodes,
scopes_by_node=scopes_by_node,
owner_stack_by_node=owner_stack_by_node,
),
items,
)
@@ -250,7 +228,7 @@ def _schema_at_path(
parts: tuple[str, ...],
active_scope: FrameScope,
foreach_nodes: Mapping[str, ForeachNode],
scopes_by_node: Mapping[str, set[FrameScope]],
owner_stack_by_node: Mapping[str, ForeachOwnerStack],
) -> Mapping[str, object] | None:
try:
schema_document = _schema_document(
@@ -258,7 +236,7 @@ def _schema_at_path(
root,
active_scope=active_scope,
foreach_nodes=foreach_nodes,
scopes_by_node=scopes_by_node,
owner_stack_by_node=owner_stack_by_node,
)
current: object = schema_document
for part in parts:
@@ -282,7 +260,7 @@ def _schema_document(
*,
active_scope: FrameScope = 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]:
if root == "input":
return workflow.input_schema.model_dump(mode="json", exclude_none=True)
@@ -295,7 +273,7 @@ def _schema_document(
if (
active_scope 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)
if foreach is not None:
@@ -307,9 +285,8 @@ def _schema_document(
_foreach_item_schema(
workflow,
foreach,
scopes_by_node.get(foreach.id, {None}),
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=[
{"from": "each", "outcome": "loop", "to": "body"},
{"from": "each", "outcome": "done", "to": "tail"},
{"from": "body", "outcome": "ok", "to": END},
{"from": "body", "outcome": "ok", "to": "each"},
{"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")],
edges=[
{"from": "each", "outcome": "loop", "to": "body"},
{"from": "body", "outcome": "ok", "to": END},
{"from": "body", "outcome": "ok", "to": "each"},
{"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")],
edges=[
{"from": "each", "outcome": "loop", "to": "body"},
{"from": "body", "outcome": "ok", "to": END},
{"from": "body", "outcome": "ok", "to": "each"},
{"from": "each", "outcome": "done", "to": END},
],
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"}}
def test_only_foreach_reachable_node_has_available_context() -> None:
def test_region_conflicted_node_receives_no_guaranteed_foreach_fields() -> None:
workflow = _workflow(
start="start",
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": "each", "outcome": "loop", "to": "body"},
{"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"
assert _field_map(workflow, "body")["item"].reason
fields = context_fields_by_node(workflow)
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:
@@ -239,8 +245,8 @@ def test_nested_foreach_replaces_inner_scope_and_restores_outer_scope() -> None:
{"from": "outer", "outcome": "loop", "to": "inner"},
{"from": "inner", "outcome": "loop", "to": "inner_body"},
{"from": "inner", "outcome": "done", "to": "after_inner"},
{"from": "inner_body", "outcome": "ok", "to": END},
{"from": "after_inner", "outcome": "ok", "to": END},
{"from": "inner_body", "outcome": "ok", "to": "inner"},
{"from": "after_inner", "outcome": "ok", "to": "outer"},
{"from": "outer", "outcome": "done", "to": END},
],
state_schema={
@@ -267,12 +273,14 @@ def test_nested_foreach_preserves_context_backed_item_schema() -> None:
_foreach("outer", alias="outer_item"),
_foreach("inner", alias="inner_item", over="context.outer_item"),
_node("inner_body"),
_node("after_inner"),
],
edges=[
{"from": "outer", "outcome": "loop", "to": "inner"},
{"from": "inner", "outcome": "loop", "to": "inner_body"},
{"from": "inner", "outcome": "done", "to": END},
{"from": "inner_body", "outcome": "ok", "to": END},
{"from": "inner", "outcome": "done", "to": "after_inner"},
{"from": "inner_body", "outcome": "ok", "to": "inner"},
{"from": "after_inner", "outcome": "ok", "to": "outer"},
{"from": "outer", "outcome": "done", "to": END},
],
state_schema={
@@ -350,4 +358,6 @@ def test_scoped_cycle_terminates_and_preserves_scoped_field_availability() -> No
fields = context_fields_by_node(workflow)
assert fields["body"]
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")