Files
lda-wf/src/wf_core/runtime/ops/state.py
T

572 lines
20 KiB
Python

from __future__ import annotations
from collections.abc import Mapping, Sequence
from copy import deepcopy
from dataclasses import dataclass
from dataclasses import field as dataclass_field
from typing import Any
from wf_core.conditions import safe_resolve_path
from wf_core.errors import WorkflowExecutionError
from wf_core.local_paths import (
LocalPathError,
get_local_value,
has_overlapping_paths,
set_local_value,
)
from wf_core.models.reducers import ReducerRef
from wf_core.models.schemas import StateFieldDecl
from wf_core.models.steps import (
InputPathBinding,
InputValueBinding,
NodeUse,
OutputBinding,
)
from wf_core.models.workflow import Workflow
from wf_core.paths import (
PathResolutionError,
StatePath,
get_nested_value,
path_parts_overlap,
set_nested_value,
split_graph_path,
)
from wf_core.run_state import StateWrite
from wf_core.runtime.ops.merges import (
ReducerDefinition,
apply_reducer,
reducer_allows_sibling_writes,
)
from wf_core.runtime.ops.schemas import validate_payload_against_schema
_MISSING = object()
@dataclass(slots=True)
class StatePatch:
"""Validated state writes produced by one step before commit.
`changes` is the public trace-facing incoming-value view. `writes` preserves
the reducer-aware records needed by lineage overlays and future gathers:
barriers replay `incoming_value`, while same-lineage reads use
`visible_value`.
"""
changes: dict[str, Any] = dataclass_field(default_factory=dict)
writes: list[StateWrite] = dataclass_field(default_factory=list)
_prepared_writes: dict[StatePath, tuple[list[str], Any]] = dataclass_field(
default_factory=dict,
repr=False,
)
_staged_state: dict[str, Any] = dataclass_field(default_factory=dict, repr=False)
def __post_init__(self) -> None:
"""Keep legacy `StatePatch(changes=...)` usable during migration.
New runtime code should prefer ordered `writes`. `changes` stays as the
public trace-facing view and as parse compatibility for old barrier
metadata/tests that predate `StateWrite`. If both are supplied, they
must describe identical incoming writes; otherwise trace and replay
semantics would disagree.
"""
if self.changes and self.writes:
derived_changes = {
str(write.path): write.incoming_value for write in self.writes
}
if self.changes != derived_changes:
raise ValueError(
"StatePatch constructed with inconsistent changes and writes"
)
return
if not self.changes and self.writes:
self.changes = {
str(write.path): write.incoming_value for write in self.writes
}
if not self.writes and self.changes:
self.writes = [
StateWrite(
path=StatePath.parse(destination),
incoming_value=value,
visible_value=value,
reducer=ReducerRef(name="wf.std.replace"),
)
for destination, value in self.changes.items()
]
@property
def visible_values(self) -> dict[str, Any]:
"""Final values visible to later reads in the same lineage."""
return {str(write.path): write.visible_value for write in self.writes}
def extend(self, patch: StatePatch) -> None:
"""Append another patch from the same lineage.
Multi-step foreach item bodies accumulate several node patches before a
barrier sees them. Both the legacy trace view and the ordered reducer
write records must be preserved.
"""
self.changes.update(patch.changes)
self.writes.extend(patch.writes)
@dataclass(slots=True, frozen=True)
class _BarrierWrite:
"""One item-lineage write observed before a barrier commit."""
item_index: int
path: StatePath
source_key: str
def apply_output_map(
workflow: Workflow,
node: NodeUse,
node_output: dict[str, Any],
state: dict[str, Any],
reducers: Mapping[str, ReducerDefinition] | None = None,
) -> dict[str, Any]:
"""Compatibility wrapper for callers that still invoke the old helper."""
try:
return apply_output_bindings(
workflow,
node.output,
node_output,
state,
reducers=reducers,
missing_field_message=(
f"node {node.id!r} did not return required mapped field {{field}}"
),
)
except AttributeError as exc:
raise WorkflowExecutionError(
"apply_output_map requires NodeUse.output canonical bindings"
) from exc
def apply_output_bindings(
workflow: Workflow,
bindings: Sequence[OutputBinding],
node_output: dict[str, Any],
state: dict[str, Any],
*,
reducers: Mapping[str, ReducerDefinition] | None = None,
missing_field_message: str = "node output did not include required field {field}",
) -> dict[str, Any]:
"""Prepare and commit one atomic state patch from canonical output bindings."""
patch = build_output_patch(
workflow,
bindings,
node_output,
state,
reducers=reducers,
missing_field_message=missing_field_message,
)
return commit_state_patch(state, patch)
def build_output_patch(
workflow: Workflow,
bindings: Sequence[OutputBinding],
node_output: Mapping[str, Any],
state: dict[str, Any],
*,
reducers: Mapping[str, ReducerDefinition] | None = None,
missing_field_message: str = "node output did not include required field {field}",
) -> StatePatch:
"""Build and validate one reducer-aware state patch without mutating state."""
if has_overlapping_paths(str(binding.target) for binding in bindings):
raise WorkflowExecutionError(
"mapped state patch has overlapping destination paths"
)
state_fields = workflow.state_schema.field_index()
resolved_patch: dict[StatePath, Any] = {}
for binding in bindings:
try:
value = get_local_value(node_output, binding.source)
except LocalPathError:
raise WorkflowExecutionError(
missing_field_message.format(field=repr(str(binding.source)))
) from None
resolved_patch[binding.target] = value
prepared_patch: dict[StatePath, tuple[list[str], Any]] = {}
for destination_path, value in resolved_patch.items():
key_path, merged_value = prepare_state_value(
workflow,
state,
destination_path,
value,
reducers=reducers,
state_fields=state_fields,
)
prepared_patch[destination_path] = (key_path, merged_value)
writes = [
StateWrite(
path=destination_path,
incoming_value=resolved_patch[destination_path],
visible_value=merged_value,
reducer=reducer_for_state_path(destination_path, state_fields),
)
for destination_path, (_key_path, merged_value) in prepared_patch.items()
]
# Stage writes on a copy so commit-time path errors cannot partially mutate state.
staged_state = deepcopy(state)
for _destination_path, (key_path, merged_value) in prepared_patch.items():
safe_set_nested_value(staged_state, key_path, merged_value)
validate_staged_state_patch(staged_state, prepared_patch, state_fields)
return StatePatch(
changes={str(path): value for path, value in resolved_patch.items()},
writes=writes,
_prepared_writes=prepared_patch,
_staged_state=staged_state,
)
def commit_state_patch(state: dict[str, Any], patch: StatePatch) -> dict[str, Any]:
"""Commit a prevalidated patch to state and return trace-facing changes."""
state.clear()
state.update(patch._staged_state)
return dict(patch.changes)
def build_barrier_patch(
workflow: Workflow,
item_patches: Sequence[StatePatch],
state: dict[str, Any],
*,
reducers: Mapping[str, ReducerDefinition] | None = None,
) -> StatePatch:
"""Build one committed barrier patch by replaying item writes in order.
Item patches are built against the parent-visible state. Their prepared
writes cannot be blindly merged because reducers must see the value produced
by earlier item patches. The barrier therefore replays trace-facing incoming
changes against a single staged state in deterministic item order.
Unlike ordinary node patches, barrier patch `changes` report the final
committed aggregate values. A barrier trace is the single visible state
commit for all buffered item patches, so showing raw per-item incoming
values would hide what actually landed in `RunState.state`.
The emitted `writes` log keeps every constituent item write in order
instead of one merged write per path. A combined patch buffered in a
lineage can itself be re-merged by an outer barrier, and replaying merged
cumulative values would duplicate whatever was already committed when the
constituents were built. Replaying the original per-item deltas stays
correct at any nesting depth. Each kept write still carries the merged
aggregate as its `visible_value`, so overlay reads and `visible_values`
keep showing the final value.
"""
state_fields = workflow.state_schema.field_index()
validate_barrier_writes(item_patches, state_fields, reducers=reducers)
staged_state = deepcopy(state)
prepared_patch: dict[StatePath, tuple[list[str], Any]] = {}
committed_changes: dict[str, Any] = {}
for item_patch in item_patches:
for write in item_patch.writes:
destination_path = write.path
key_path, merged_value = prepare_state_value(
workflow,
staged_state,
destination_path,
write.incoming_value,
reducers=reducers,
state_fields=state_fields,
)
safe_set_nested_value(staged_state, key_path, merged_value)
prepared_patch[destination_path] = (key_path, merged_value)
committed_changes[str(destination_path)] = merged_value
merged_visible = {
destination_path: merged_value
for destination_path, (_key_path, merged_value) in prepared_patch.items()
}
writes = [
StateWrite(
path=write.path,
incoming_value=write.incoming_value,
visible_value=merged_visible[write.path],
reducer=write.reducer,
)
for item_patch in item_patches
for write in item_patch.writes
]
validate_staged_state_patch(staged_state, prepared_patch, state_fields)
combined = StatePatch(
writes=writes,
_prepared_writes=prepared_patch,
_staged_state=staged_state,
)
# The trace-facing view reports the aggregate, while the replay log above
# intentionally carries per-item deltas (see docstring). Assign it after
# construction: passing both to the constructor requires them to agree.
combined.changes = committed_changes
return combined
def validate_barrier_writes(
item_patches: Sequence[StatePatch],
state_fields: Mapping[StatePath, StateFieldDecl],
*,
reducers: Mapping[str, ReducerDefinition] | None = None,
) -> None:
"""Reject ambiguous sibling writes before replaying a foreach barrier.
Normal node patch validation handles one node output. This helper handles
writes from different foreach item lineages that commit together.
"""
writes = _barrier_writes(item_patches)
for index, left in enumerate(writes):
for right in writes[index + 1 :]:
if left.item_index == right.item_index:
continue
if left.path == right.path:
if _allows_sibling_writes(left.path, state_fields, reducers):
continue
raise WorkflowExecutionError(
"multiple sibling writes to "
f"{left.source_key!r} require a mergeable reducer"
)
if _state_paths_overlap(left.path, right.path):
raise WorkflowExecutionError(
"overlapping sibling writes are not supported at a foreach "
f"barrier: {left.source_key!r} and {right.source_key!r}"
)
def _barrier_writes(item_patches: Sequence[StatePatch]) -> list[_BarrierWrite]:
writes: list[_BarrierWrite] = []
for item_index, item_patch in enumerate(item_patches):
for destination in item_patch.changes:
path = StatePath.parse(destination)
writes.append(
_BarrierWrite(
item_index=item_index,
path=path,
source_key=destination,
)
)
return writes
def _allows_sibling_writes(
path: StatePath,
state_fields: Mapping[StatePath, StateFieldDecl],
reducers: Mapping[str, ReducerDefinition] | None,
) -> bool:
field = state_fields.get(path)
if field is None or field.reducer is None:
return False
return reducer_allows_sibling_writes(field.reducer, reducers)
def _state_paths_overlap(left: StatePath, right: StatePath) -> bool:
return path_parts_overlap(left.parts, right.parts)
def apply_mapped_state(
workflow: Workflow,
source_data: dict[str, Any],
mapping: dict[str, str],
state: dict[str, Any],
*,
reducers: Mapping[str, ReducerDefinition] | None = None,
missing_field_message: str,
) -> dict[str, Any]:
bindings = [
OutputBinding.model_validate({"source": source, "target": target})
for source, target in mapping.items()
]
return apply_output_bindings(
workflow,
bindings,
source_data,
state,
reducers=reducers,
missing_field_message=missing_field_message,
)
def write_state_value(
workflow: Workflow,
state: dict[str, Any],
destination_path: str,
value: Any,
*,
reducers: Mapping[str, ReducerDefinition] | None = None,
) -> None:
key_path, merged_value = prepare_state_value(
workflow,
state,
destination_path,
value,
reducers=reducers,
)
staged_state = deepcopy(state)
safe_set_nested_value(staged_state, key_path, merged_value)
validate_staged_state_patch(
staged_state,
{StatePath.parse(destination_path): (key_path, merged_value)},
workflow.state_schema.field_index(),
)
state.clear()
state.update(staged_state)
def prepare_state_value(
workflow: Workflow,
state: dict[str, Any],
destination_path: str | StatePath,
value: Any,
*,
reducers: Mapping[str, ReducerDefinition] | None = None,
state_fields: Mapping[StatePath, StateFieldDecl] | None = None,
) -> tuple[list[str], Any]:
"""Resolve reducer output for a state write without mutating state."""
try:
root, parts = split_graph_path(destination_path)
except PathResolutionError as exc:
raise WorkflowExecutionError(str(exc)) from exc
if root != "state":
raise WorkflowExecutionError(
f"executor only supports writes into state.*, got {destination_path!r}"
)
declared_path = StatePath(tuple(parts))
fields = (
state_fields
if state_fields is not None
else workflow.state_schema.field_index()
)
declared_field = fields.get(declared_path)
reducer = (
declared_field.reducer if declared_field else ReducerRef(name="wf.std.replace")
)
key_path = parts
current_value = get_nested_value(state, key_path)
merged_value = apply_reducer(
reducer=reducer,
current_value=current_value,
incoming_value=value,
destination_path=str(destination_path),
reducers=reducers,
)
return key_path, merged_value
def reducer_for_state_path(
path: StatePath,
state_fields: Mapping[StatePath, StateFieldDecl],
) -> ReducerRef:
"""Return the reducer declared for one exact state path."""
declared_field = state_fields.get(path)
return (
declared_field.reducer if declared_field else ReducerRef(name="wf.std.replace")
)
def project_output(
workflow: Workflow,
state: dict[str, Any],
*,
workflow_input: Mapping[str, Any] | None = None,
context: Mapping[str, Any] | None = None,
) -> dict[str, Any]:
"""Project final workflow output from explicit bindings or state fields.
`workflow.output` is the canonical root-output mapping for workflows whose
public output shape does not mirror top-level state keys. Older workflows
without explicit output bindings keep the same-name state projection.
"""
if workflow.output:
output: dict[str, Any] = {}
for binding in workflow.output:
if isinstance(binding, InputValueBinding):
value = binding.value
elif isinstance(binding, InputPathBinding):
value = safe_resolve_path(
str(binding.path),
state=state,
workflow_input=workflow_input or {},
context=context or {},
)
else:
raise WorkflowExecutionError("unsupported workflow output binding")
try:
set_local_value(output, binding.target, value)
except LocalPathError as exc:
raise WorkflowExecutionError(str(exc)) from exc
return output
return {
key: state[key] for key in workflow.output_schema.properties if key in state
}
def validate_staged_state_patch(
staged_state: dict[str, Any],
prepared_patch: Mapping[StatePath, tuple[list[str], Any]],
state_fields: Mapping[StatePath, StateFieldDecl],
) -> None:
"""Validate affected declared state schemas before committing a patch.
Runtime writes are path-based, while JSON Schema is tree-based. A child
write can violate a declared parent schema, and a parent replacement can
violate declared child schemas. This helper validates every declared schema
that is on either side of a staged write path, without mutating the original
state first.
"""
for field in _affected_state_fields(prepared_patch, state_fields):
value = _get_existing_nested_value(staged_state, list(field.path.parts))
if value is _MISSING:
continue
validate_payload_against_schema(
field.validation_schema,
value,
f"state write state.{'.'.join(field.path.parts)}",
)
def _affected_state_fields(
prepared_patch: Mapping[StatePath, tuple[list[str], Any]],
state_fields: Mapping[StatePath, StateFieldDecl],
) -> list[StateFieldDecl]:
affected: dict[StatePath, StateFieldDecl] = {}
for destination_path in prepared_patch:
destination_parts = destination_path.parts
for path, field in state_fields.items():
field_parts = field.path.parts
if _is_prefix(field_parts, destination_parts) or _is_prefix(
destination_parts,
field_parts,
):
affected[path] = field
return sorted(
affected.values(),
key=lambda field: len(field.path.parts),
reverse=True,
)
def _is_prefix(prefix: tuple[str, ...], value: tuple[str, ...]) -> bool:
return len(prefix) <= len(value) and value[: len(prefix)] == prefix
def _get_existing_nested_value(state: Mapping[str, Any], path_parts: list[str]) -> Any:
current: Any = state
for part in path_parts:
if not isinstance(current, Mapping) or part not in current:
return _MISSING
current = current[part]
return current
def safe_set_nested_value(
state: dict[str, Any], path_parts: list[str], value: Any
) -> None:
try:
set_nested_value(state, path_parts, value)
except PathResolutionError as exc:
raise WorkflowExecutionError(str(exc)) from exc