typed path, path handling,

half the plan is done
This commit is contained in:
lda
2026-05-20 16:05:43 +07:00 Verified
parent a74e016f7b
commit 3683c23937
11 changed files with 1316 additions and 189 deletions
+24 -18
View File
@@ -3,36 +3,42 @@ from __future__ import annotations
from collections.abc import Iterable, Mapping from collections.abc import Iterable, Mapping
from typing import Any from typing import Any
from wf_core.paths import LocalPath, PathResolutionError
class LocalPathError(ValueError): class LocalPathError(ValueError):
"""Raised when a node-local dotted path cannot be parsed or resolved.""" """Raised when a node-local dotted path cannot be parsed or resolved."""
def split_local_path(path: str) -> list[str]: def _coerce_local_path(path: str | LocalPath) -> LocalPath:
"""Split one dotted node-local path, rejecting empty segments.""" try:
if path == ".": return path if isinstance(path, LocalPath) else LocalPath.parse(path)
return [] except PathResolutionError as exc:
parts = path.split(".") raise LocalPathError(str(exc)) from exc
if not path or any(not part for part in parts):
raise LocalPathError(f"invalid local path {path!r}")
return parts
def get_local_value(payload: Mapping[str, Any], path: str) -> Any: def split_local_path(path: str | LocalPath) -> list[str]:
"""Split one node-local path, accepting the new typed path object."""
return list(_coerce_local_path(path).parts)
def get_local_value(payload: Mapping[str, Any], path: str | LocalPath) -> Any:
"""Resolve one node-local path from a nested mapping payload.""" """Resolve one node-local path from a nested mapping payload."""
if path == ".": parsed = _coerce_local_path(path)
if not parsed.parts:
return dict(payload) return dict(payload)
current: Any = payload current: Any = payload
for part in split_local_path(path): for part in parsed.parts:
if not isinstance(current, Mapping) or part not in current: if not isinstance(current, Mapping) or part not in current:
raise LocalPathError(f"local path {path!r} could not be resolved") raise LocalPathError(f"local path {str(parsed)!r} could not be resolved")
current = current[part] current = current[part]
return current return current
def set_local_value(payload: dict[str, Any], path: str, value: Any) -> None: def set_local_value(payload: dict[str, Any], path: str | LocalPath, value: Any) -> None:
"""Write one value into a nested node-local mapping payload.""" """Write one value into a nested node-local mapping payload."""
parts = split_local_path(path) parsed = _coerce_local_path(path)
parts = list(parsed.parts)
if not parts: if not parts:
if not isinstance(value, Mapping): if not isinstance(value, Mapping):
raise LocalPathError("root local path requires a mapping value") raise LocalPathError("root local path requires a mapping value")
@@ -43,12 +49,12 @@ def set_local_value(payload: dict[str, Any], path: str, value: Any) -> None:
for part in parts[:-1]: for part in parts[:-1]:
next_value = current.setdefault(part, {}) next_value = current.setdefault(part, {})
if not isinstance(next_value, dict): if not isinstance(next_value, dict):
raise LocalPathError(f"local path {path!r} overlaps an existing value") raise LocalPathError(f"local path {str(parsed)!r} overlaps an existing value")
current = next_value current = next_value
current[parts[-1]] = value current[parts[-1]] = value
def paths_overlap(left: str, right: str) -> bool: def paths_overlap(left: str | LocalPath, right: str | LocalPath) -> bool:
"""Return whether two dotted paths overlap by equality or ancestry.""" """Return whether two dotted paths overlap by equality or ancestry."""
left_parts = split_local_path(left) left_parts = split_local_path(left)
right_parts = split_local_path(right) right_parts = split_local_path(right)
@@ -56,9 +62,9 @@ def paths_overlap(left: str, right: str) -> bool:
return left_parts[:shortest] == right_parts[:shortest] return left_parts[:shortest] == right_parts[:shortest]
def has_overlapping_paths(paths: Iterable[str]) -> bool: def has_overlapping_paths(paths: Iterable[str | LocalPath]) -> bool:
"""Return whether any pair of dotted paths overlaps.""" """Return whether any pair of dotted paths overlaps."""
seen: list[str] = [] seen: list[str | LocalPath] = []
for path in paths: for path in paths:
if any(paths_overlap(path, prior) for prior in seen): if any(paths_overlap(path, prior) for prior in seen):
return True return True
+91 -24
View File
@@ -1,10 +1,46 @@
from __future__ import annotations from __future__ import annotations
from collections.abc import Mapping
from typing import Annotated, Literal from typing import Annotated, Literal
from pydantic import BaseModel, ConfigDict, Field from pydantic import BaseModel, ConfigDict, Field, model_validator
from wf_core.models.conditions import Condition from wf_core.models.conditions import Condition
from wf_core.paths import GraphSourcePath, LocalPath, StatePath
class InputPathBinding(BaseModel):
"""Map one workflow graph source path into one node-local input path."""
model_config = ConfigDict(extra="forbid")
target: LocalPath
path: GraphSourcePath
class InputValueBinding(BaseModel):
"""Map one static value into one node-local input path."""
model_config = ConfigDict(extra="forbid")
target: LocalPath
value: object
InputBinding = Annotated[
InputPathBinding | InputValueBinding,
Field(union_mode="left_to_right"),
]
"""Canonical node input binding, distinguished by `path` vs `value` shape."""
class OutputBinding(BaseModel):
"""Map one node-local output path into one workflow state path."""
model_config = ConfigDict(extra="forbid")
source: LocalPath
target: StatePath
class NodeUse(BaseModel): class NodeUse(BaseModel):
@@ -14,32 +50,63 @@ class NodeUse(BaseModel):
type: Literal["node"] type: Literal["node"]
node: str node: str
desc: str | None = None desc: str | None = None
in_map: dict[str, str] = Field( input: list[InputBinding] = Field(default_factory=list)
default_factory=dict, output: list[OutputBinding] = Field(default_factory=list)
description=(
"Map graph source paths to node-local input paths. Keys are paths "
"such as input.text, state.user.name, or context.item; values are "
"input fields/paths inside the node payload."
),
)
input_values: dict[str, object] = Field(
default_factory=dict,
description=(
"Static node-local input values keyed by destination input field/path. "
"Use this for graph-defined constants; use in_map only for graph paths."
),
)
out_map: dict[str, str] = Field(
default_factory=dict,
description=(
"Map node-local output paths to workflow state destinations. Keys "
"are output fields/paths inside the node payload; values must be "
"state.* destination paths."
),
)
retry: int | None = Field(default=None, ge=0) retry: int | None = Field(default=None, ge=0)
timeout_seconds: int | None = Field(default=None, gt=0) timeout_seconds: int | None = Field(default=None, gt=0)
@model_validator(mode="before")
@classmethod
def _coerce_deprecated_maps(cls, data: object) -> object:
"""Normalize deprecated map fields into canonical parse-only bindings."""
if not isinstance(data, Mapping):
return data
old_fields = ("in_map", "input_values", "out_map")
has_canonical = "input" in data or "output" in data
present_old_fields = [field for field in old_fields if field in data]
if has_canonical and present_old_fields:
old_names = ", ".join(present_old_fields)
raise ValueError(
f"cannot mix canonical input/output with deprecated fields: {old_names}"
)
normalized = dict(data)
input_bindings = list(normalized.pop("input", []))
output_bindings = list(normalized.pop("output", []))
input_values = cls._deprecated_mapping(
normalized.pop("input_values", {}), field_name="input_values"
)
in_map = cls._deprecated_mapping(normalized.pop("in_map", {}), field_name="in_map")
out_map = cls._deprecated_mapping(
normalized.pop("out_map", {}), field_name="out_map"
)
input_bindings.extend(
{"target": target, "value": value}
for target, value in input_values.items()
)
input_bindings.extend(
{"target": target, "path": path} for path, target in in_map.items()
)
output_bindings.extend(
{"source": source, "target": target} for source, target in out_map.items()
)
normalized["input"] = input_bindings
normalized["output"] = output_bindings
return normalized
@staticmethod
def _deprecated_mapping(
value: object, *, field_name: str
) -> Mapping[object, object]:
"""Reject malformed deprecated map inputs before calling `.items()`."""
if not isinstance(value, Mapping):
raise ValueError(f"{field_name} must be a mapping")
return value
class ConditionNode(BaseModel): class ConditionNode(BaseModel):
"""Control-flow step that routes through `true` or `false` outcomes.""" """Control-flow step that routes through `true` or `false` outcomes."""
+246 -11
View File
@@ -1,22 +1,252 @@
from __future__ import annotations from __future__ import annotations
from dataclasses import dataclass
import re
from collections.abc import Mapping, MutableMapping from collections.abc import Mapping, MutableMapping
from typing import Any from typing import Any, ClassVar, Literal
from pydantic_core import core_schema
class PathResolutionError(ValueError): class PathResolutionError(ValueError):
pass pass
def split_graph_path(path: str) -> tuple[str, list[str]]: SEGMENT_RE = re.compile(r"^[A-Za-z_][A-Za-z0-9_]*$")
root, *parts = path.split(".") GraphRoot = Literal["input", "state", "context"]
if not root or not parts:
raise PathResolutionError(f"invalid path {path!r}")
def _validate_segment(segment: str, *, path_kind: str) -> str:
if not SEGMENT_RE.fullmatch(segment):
raise PathResolutionError(f"invalid {path_kind} segment {segment!r}")
return segment
def _parse_fragments(*fragments: str, path_kind: str) -> tuple[str, ...]:
parts: list[str] = []
for fragment in fragments:
if not fragment:
raise PathResolutionError(f"invalid {path_kind} path")
fragment_parts = fragment.split(".")
if any(not part for part in fragment_parts):
raise PathResolutionError(f"invalid {path_kind} path {fragment!r}")
parts.extend(
_validate_segment(part, path_kind=path_kind) for part in fragment_parts
)
return tuple(parts)
def _json_schema(pattern: str, description: str) -> dict[str, Any]:
return {"type": "string", "pattern": pattern, "description": description}
@dataclass(frozen=True)
class LocalPath:
"""Node-local payload path. The root marker `.` means the whole payload."""
parts: tuple[str, ...]
_JSON_PATTERN: ClassVar[str] = (
r"^(\.|[A-Za-z_][A-Za-z0-9_]*(\.[A-Za-z_][A-Za-z0-9_]*)*)$"
)
def __post_init__(self) -> None:
object.__setattr__(
self,
"parts",
tuple(_validate_segment(part, path_kind="local") for part in self.parts),
)
@classmethod
def root(cls) -> LocalPath:
return cls(())
@classmethod
def of(cls, *fragments: str) -> LocalPath:
if not fragments:
return cls.root()
return cls(_parse_fragments(*fragments, path_kind="local"))
@classmethod
def parse(cls, raw: str) -> LocalPath:
if raw == ".":
return cls.root()
return cls.of(raw)
def __str__(self) -> str:
return "." if not self.parts else ".".join(self.parts)
@classmethod
def __get_pydantic_core_schema__(
cls, _source_type: object, _handler: object
) -> core_schema.CoreSchema:
def validate(value: object) -> LocalPath:
if isinstance(value, cls):
return cls(value.parts)
if isinstance(value, str):
return cls.parse(value)
raise ValueError("expected local path string")
return core_schema.no_info_plain_validator_function(
validate,
serialization=core_schema.plain_serializer_function_ser_schema(
str, when_used="json"
),
)
@classmethod
def __get_pydantic_json_schema__(
cls, _core_schema: core_schema.CoreSchema, _handler: object
) -> dict[str, Any]:
return _json_schema(
cls._JSON_PATTERN,
"Node-local dotted path or root marker `.`.",
)
@dataclass(frozen=True)
class GraphSourcePath:
"""Readable workflow graph path rooted at input, state, or context."""
root: GraphRoot
parts: tuple[str, ...] = ()
_ROOTS: ClassVar[set[str]] = {"input", "state", "context"}
_JSON_PATTERN: ClassVar[str] = (
r"^(input|state|context)(\.[A-Za-z_][A-Za-z0-9_]*)*$"
)
def __post_init__(self) -> None:
if self.root not in self._ROOTS:
raise PathResolutionError(f"unknown path root {self.root!r}")
object.__setattr__(
self,
"parts",
tuple(
_validate_segment(part, path_kind="graph source")
for part in self.parts
),
)
@classmethod
def parse(cls, raw: str) -> GraphSourcePath:
root, *raw_parts = raw.split(".")
if root not in cls._ROOTS:
raise PathResolutionError(f"unknown path root {root!r}")
parts = tuple(
_validate_segment(part, path_kind="graph source") for part in raw_parts
)
return cls(root, parts) # type: ignore[arg-type]
@classmethod
def input(cls, *fragments: str) -> GraphSourcePath:
return cls("input", _parse_fragments(*fragments, path_kind="graph source"))
@classmethod
def state(cls, *fragments: str) -> GraphSourcePath:
return cls("state", _parse_fragments(*fragments, path_kind="graph source"))
@classmethod
def context(cls, *fragments: str) -> GraphSourcePath:
return cls("context", _parse_fragments(*fragments, path_kind="graph source"))
def __str__(self) -> str:
return self.root if not self.parts else f"{self.root}.{'.'.join(self.parts)}"
@classmethod
def __get_pydantic_core_schema__(
cls, _source_type: object, _handler: object
) -> core_schema.CoreSchema:
def validate(value: object) -> GraphSourcePath:
if isinstance(value, cls):
return cls(value.root, value.parts)
if isinstance(value, str):
return cls.parse(value)
raise ValueError("expected graph source path string")
return core_schema.no_info_plain_validator_function(
validate,
serialization=core_schema.plain_serializer_function_ser_schema(
str, when_used="json"
),
)
@classmethod
def __get_pydantic_json_schema__(
cls, _core_schema: core_schema.CoreSchema, _handler: object
) -> dict[str, Any]:
return _json_schema(
cls._JSON_PATTERN,
"Readable graph path rooted at input, state, or context.",
)
@dataclass(frozen=True)
class StatePath:
"""Writable workflow state path. Bare `state` is intentionally invalid."""
parts: tuple[str, ...]
_JSON_PATTERN: ClassVar[str] = r"^state\.[A-Za-z_][A-Za-z0-9_]*(\.[A-Za-z_][A-Za-z0-9_]*)*$"
def __post_init__(self) -> None:
parts = tuple(_validate_segment(part, path_kind="state") for part in self.parts)
if not parts:
raise PathResolutionError("expected state path such as state.foo")
object.__setattr__(self, "parts", parts)
@classmethod
def of(cls, *fragments: str) -> StatePath:
parts = _parse_fragments(*fragments, path_kind="state")
if not parts:
raise PathResolutionError("expected state path such as state.foo")
return cls(parts)
@classmethod
def parse(cls, raw: str) -> StatePath:
parsed = GraphSourcePath.parse(raw)
if parsed.root != "state" or not parsed.parts:
raise PathResolutionError("expected state path such as state.foo")
return cls(parsed.parts)
def __str__(self) -> str:
return f"state.{'.'.join(self.parts)}"
@classmethod
def __get_pydantic_core_schema__(
cls, _source_type: object, _handler: object
) -> core_schema.CoreSchema:
def validate(value: object) -> StatePath:
if isinstance(value, cls):
return cls(value.parts)
if isinstance(value, str):
return cls.parse(value)
raise ValueError("expected state path string")
return core_schema.no_info_plain_validator_function(
validate,
serialization=core_schema.plain_serializer_function_ser_schema(
str, when_used="json"
),
)
@classmethod
def __get_pydantic_json_schema__(
cls, _core_schema: core_schema.CoreSchema, _handler: object
) -> dict[str, Any]:
return _json_schema(cls._JSON_PATTERN, "Writable state path such as state.foo.")
def split_graph_path(path: str | GraphSourcePath | StatePath) -> tuple[str, list[str]]:
if isinstance(path, StatePath):
return "state", list(path.parts)
parsed = path if isinstance(path, GraphSourcePath) else GraphSourcePath.parse(path)
root, parts = parsed.root, list(parsed.parts)
return root, parts return root, parts
def is_valid_source_path( def is_valid_source_path(
path: str, path: str | GraphSourcePath,
state_root_fields: set[str], state_root_fields: set[str],
input_root_fields: set[str], input_root_fields: set[str],
*, *,
@@ -27,6 +257,11 @@ def is_valid_source_path(
except PathResolutionError: except PathResolutionError:
return False return False
if not parts:
return root in {"input", "state", "context"} and (
root != "context" or allow_context
)
field_name = parts[0] field_name = parts[0]
if allow_context and root == "context": if allow_context and root == "context":
return True return True
@@ -37,16 +272,16 @@ def is_valid_source_path(
return False return False
def is_valid_destination_path(path: str) -> bool: def is_valid_destination_path(path: str | StatePath) -> bool:
try: try:
root, parts = split_graph_path(path) StatePath.parse(str(path))
except PathResolutionError: except PathResolutionError:
return False return False
return root == "state" and bool(parts) return True
def resolve_graph_path( def resolve_graph_path(
path: str, path: str | GraphSourcePath,
*, *,
state: Mapping[str, Any], state: Mapping[str, Any],
workflow_input: Mapping[str, Any], workflow_input: Mapping[str, Any],
@@ -72,7 +307,7 @@ def resolve_graph_path(
def path_exists( def path_exists(
path: str, path: str | GraphSourcePath,
*, *,
state: Mapping[str, Any], state: Mapping[str, Any],
workflow_input: Mapping[str, Any], workflow_input: Mapping[str, Any],
+14 -12
View File
@@ -8,13 +8,13 @@ from wf_core.errors import WorkflowExecutionError
from wf_core.local_paths import LocalPathError, set_local_value from wf_core.local_paths import LocalPathError, set_local_value
from wf_core.models.results import NodeResult from wf_core.models.results import NodeResult
from wf_core.models.schemas import NodeDef from wf_core.models.schemas import NodeDef
from wf_core.models.steps import NodeUse from wf_core.models.steps import InputPathBinding, InputValueBinding, NodeUse
from wf_core.models.workflow import Workflow from wf_core.models.workflow import Workflow
from wf_core.run_state import RunState, RuntimeContext, StepExecutionResult from wf_core.run_state import RunState, RuntimeContext, StepExecutionResult
from wf_core.runtime.ops.frames import frame_context_values from wf_core.runtime.ops.frames import frame_context_values
from wf_core.runtime.ops.merges import ReducerDefinition from wf_core.runtime.ops.merges import ReducerDefinition
from wf_core.runtime.ops.schemas import validate_payload_against_schema from wf_core.runtime.ops.schemas import validate_payload_against_schema
from wf_core.runtime.ops.state import apply_output_map from wf_core.runtime.ops.state import apply_output_bindings
NodeHandler = Callable[[dict[str, Any], RuntimeContext], NodeResult | dict[str, Any]] NodeHandler = Callable[[dict[str, Any], RuntimeContext], NodeResult | dict[str, Any]]
AsyncNodeHandler = Callable[ AsyncNodeHandler = Callable[
@@ -33,20 +33,22 @@ def _resolve_node_execution(
frame = run.current_frame() frame = run.current_frame()
context_values = frame_context_values(frame) context_values = frame_context_values(frame)
resolved_input: dict[str, Any] = {} resolved_input: dict[str, Any] = {}
for destination_field, value in node.input_values.items(): for binding in node.input:
try: if isinstance(binding, InputValueBinding):
set_local_value(resolved_input, destination_field, value) value = binding.value
except LocalPathError as exc: elif isinstance(binding, InputPathBinding):
raise WorkflowExecutionError(str(exc)) from exc
for source_path, destination_field in node.in_map.items():
value = safe_resolve_path( value = safe_resolve_path(
source_path, str(binding.path),
state=run.state, state=run.state,
workflow_input=run.workflow_input, workflow_input=run.workflow_input,
context=context_values, context=context_values,
) )
else:
raise WorkflowExecutionError(
f"unsupported input binding for node {node.id!r}"
)
try: try:
set_local_value(resolved_input, destination_field, value) set_local_value(resolved_input, binding.target, value)
except LocalPathError as exc: except LocalPathError as exc:
raise WorkflowExecutionError(str(exc)) from exc raise WorkflowExecutionError(str(exc)) from exc
validate_payload_against_schema( validate_payload_against_schema(
@@ -83,9 +85,9 @@ def _finalize_node_execution(
validate_payload_against_schema( validate_payload_against_schema(
node_def.output_schema, result.output, f"node output for {node.id}" node_def.output_schema, result.output, f"node output for {node.id}"
) )
state_changes = apply_output_map( state_changes = apply_output_bindings(
workflow, workflow,
node, node.output,
result.output, result.output,
run.state, run.state,
reducers=reducers, reducers=reducers,
+89 -27
View File
@@ -1,15 +1,17 @@
from __future__ import annotations from __future__ import annotations
from collections.abc import Mapping from collections.abc import Mapping, Sequence
from copy import deepcopy
from typing import Any from typing import Any
from wf_core.errors import WorkflowExecutionError from wf_core.errors import WorkflowExecutionError
from wf_core.local_paths import LocalPathError, get_local_value, has_overlapping_paths from wf_core.local_paths import LocalPathError, get_local_value, has_overlapping_paths
from wf_core.models.reducers import ReducerRef from wf_core.models.reducers import ReducerRef
from wf_core.models.steps import NodeUse from wf_core.models.steps import NodeUse, OutputBinding
from wf_core.models.workflow import Workflow from wf_core.models.workflow import Workflow
from wf_core.paths import ( from wf_core.paths import (
PathResolutionError, PathResolutionError,
StatePath,
get_nested_value, get_nested_value,
set_nested_value, set_nested_value,
split_graph_path, split_graph_path,
@@ -24,14 +26,67 @@ def apply_output_map(
state: dict[str, Any], state: dict[str, Any],
reducers: Mapping[str, ReducerDefinition] | None = None, reducers: Mapping[str, ReducerDefinition] | None = None,
) -> dict[str, Any]: ) -> dict[str, Any]:
return apply_mapped_state( """Compatibility wrapper for callers that still invoke the old helper."""
try:
return apply_output_bindings(
workflow, workflow,
node.output,
node_output, node_output,
node.out_map,
state, state,
reducers=reducers, reducers=reducers,
missing_field_message=f"node {node.id!r} did not return required mapped field {{field}}", 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."""
if has_overlapping_paths(str(binding.target) for binding in bindings):
raise WorkflowExecutionError(
"mapped state patch has overlapping destination paths"
)
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,
)
prepared_patch[destination_path] = (key_path, merged_value)
# 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)
state.clear()
state.update(staged_state)
return {str(path): value for path, value in resolved_patch.items()}
def apply_mapped_state( def apply_mapped_state(
@@ -43,30 +98,18 @@ def apply_mapped_state(
reducers: Mapping[str, ReducerDefinition] | None = None, reducers: Mapping[str, ReducerDefinition] | None = None,
missing_field_message: str, missing_field_message: str,
) -> dict[str, Any]: ) -> dict[str, Any]:
if has_overlapping_paths(mapping.values()): bindings = [
raise WorkflowExecutionError( OutputBinding.model_validate({"source": source, "target": target})
"mapped state patch has overlapping destination paths" for source, target in mapping.items()
) ]
return apply_output_bindings(
patch: dict[str, Any] = {}
for source_field, destination_path in mapping.items():
try:
value = get_local_value(source_data, source_field)
except LocalPathError:
raise WorkflowExecutionError(
missing_field_message.format(field=repr(source_field))
) from None
patch[destination_path] = value
for destination_path, value in patch.items():
write_state_value(
workflow, workflow,
bindings,
source_data,
state, state,
destination_path,
value,
reducers=reducers, reducers=reducers,
missing_field_message=missing_field_message,
) )
return dict(patch)
def write_state_value( def write_state_value(
@@ -77,6 +120,25 @@ def write_state_value(
*, *,
reducers: Mapping[str, ReducerDefinition] | None = None, reducers: Mapping[str, ReducerDefinition] | None = None,
) -> None: ) -> None:
key_path, merged_value = prepare_state_value(
workflow,
state,
destination_path,
value,
reducers=reducers,
)
safe_set_nested_value(state, key_path, merged_value)
def prepare_state_value(
workflow: Workflow,
state: dict[str, Any],
destination_path: str | StatePath,
value: Any,
*,
reducers: Mapping[str, ReducerDefinition] | None = None,
) -> tuple[list[str], Any]:
"""Resolve reducer output for a state write without mutating state."""
try: try:
root, parts = split_graph_path(destination_path) root, parts = split_graph_path(destination_path)
except PathResolutionError as exc: except PathResolutionError as exc:
@@ -98,10 +160,10 @@ def write_state_value(
reducer=reducer, reducer=reducer,
current_value=current_value, current_value=current_value,
incoming_value=value, incoming_value=value,
destination_path=destination_path, destination_path=str(destination_path),
reducers=reducers, reducers=reducers,
) )
safe_set_nested_value(state, key_path, merged_value) return key_path, merged_value
def project_output(workflow: Workflow, state: dict[str, Any]) -> dict[str, Any]: def project_output(workflow: Workflow, state: dict[str, Any]) -> dict[str, Any]:
+48 -39
View File
@@ -11,9 +11,21 @@ from wf_core.models.conditions import (
) )
from wf_core.local_paths import LocalPathError, has_overlapping_paths, split_local_path from wf_core.local_paths import LocalPathError, has_overlapping_paths, split_local_path
from wf_core.models.schemas import NodeDef from wf_core.models.schemas import NodeDef
from wf_core.models.steps import ConditionNode, ForeachNode, InterruptNode, NodeUse from wf_core.models.steps import (
ConditionNode,
ForeachNode,
InputPathBinding,
InterruptNode,
NodeUse,
)
from wf_core.models.workflow import Workflow from wf_core.models.workflow import Workflow
from wf_core.paths import is_valid_destination_path, is_valid_source_path from wf_core.paths import (
LocalPath,
PathResolutionError,
StatePath,
is_valid_destination_path,
is_valid_source_path,
)
from wf_core.validation.issues import ValidationIssueCode, ValidationReport from wf_core.validation.issues import ValidationIssueCode, ValidationReport
@@ -36,76 +48,66 @@ def validate_node_use(
input_fields = set(node_def.input_schema.properties) input_fields = set(node_def.input_schema.properties)
output_fields = set(node_def.output_schema.properties) output_fields = set(node_def.output_schema.properties)
state_fields = set(workflow.state_schema.fields) state_fields = set(workflow.state_schema.fields)
state_root_fields = {field.split(".", maxsplit=1)[0] for field in state_fields}
input_root_fields = set(workflow.input_schema.properties) input_root_fields = set(workflow.input_schema.properties)
for destination_field in node.input_values: input_targets = []
destination_root = _local_root(destination_field) for input_index, binding in enumerate(node.input):
input_targets.append(binding.target)
destination_root = _local_root(binding.target)
if destination_root is None or ( if destination_root is None or (
destination_root != "." and destination_root not in input_fields destination_root != "." and destination_root not in input_fields
): ):
report.add( report.add(
ValidationIssueCode.INVALID_NODE_INPUT_FIELD, ValidationIssueCode.INVALID_NODE_INPUT_FIELD,
f"nodes[{index}].input_values[{destination_field!r}]", f"nodes[{index}].input[{input_index}].target",
f"destination field {destination_field!r} is not declared in node input schema", f"destination field {str(binding.target)!r} is not declared in node input schema",
) )
for source_path, destination_field in node.in_map.items(): if isinstance(binding, InputPathBinding) and not is_valid_source_path(
destination_root = _local_root(destination_field) binding.path, state_root_fields, input_root_fields, allow_context=True
if destination_root is None or (
destination_root != "." and destination_root not in input_fields
):
report.add(
ValidationIssueCode.INVALID_NODE_INPUT_FIELD,
f"nodes[{index}].in_map[{source_path!r}]",
f"destination field {destination_field!r} is not declared in node input schema",
)
if not is_valid_source_path(
source_path, state_fields, input_root_fields, allow_context=True
): ):
report.add( report.add(
ValidationIssueCode.INVALID_SOURCE_PATH, ValidationIssueCode.INVALID_SOURCE_PATH,
f"nodes[{index}].in_map[{source_path!r}]", f"nodes[{index}].input[{input_index}].path",
"source path must start with input., state., or context. and reference a declared root field when applicable", "source path must start with input., state., or context. and reference a declared root field when applicable",
) )
if has_overlapping_paths(node.in_map.values()): if has_overlapping_paths(input_targets):
report.add( report.add(
ValidationIssueCode.INVALID_NODE_INPUT_FIELD, ValidationIssueCode.INVALID_NODE_INPUT_FIELD,
f"nodes[{index}].in_map", f"nodes[{index}].input",
"in_map has overlapping node-local input paths", "input has overlapping node-local input paths",
)
if has_overlapping_paths([*node.input_values, *node.in_map.values()]):
report.add(
ValidationIssueCode.INVALID_NODE_INPUT_FIELD,
f"nodes[{index}].input_values",
"static input_values overlap with path-based in_map destinations",
) )
for source_field, destination_path in node.out_map.items(): output_targets = []
source_root = _local_root(source_field) for output_index, binding in enumerate(node.output):
output_targets.append(str(binding.target))
source_root = _local_root(binding.source)
if source_root is None or ( if source_root is None or (
source_root != "." and source_root not in output_fields source_root != "." and source_root not in output_fields
): ):
report.add( report.add(
ValidationIssueCode.INVALID_NODE_OUTPUT_FIELD, ValidationIssueCode.INVALID_NODE_OUTPUT_FIELD,
f"nodes[{index}].out_map[{source_field!r}]", f"nodes[{index}].output[{output_index}].source",
f"source field {source_field!r} is not declared in node output schema", f"source field {str(binding.source)!r} is not declared in node output schema",
) )
if not is_valid_destination_path(destination_path): destination_root = _state_destination_root(binding.target)
if destination_root is None or destination_root not in state_root_fields:
report.add( report.add(
ValidationIssueCode.INVALID_DESTINATION_PATH, ValidationIssueCode.INVALID_DESTINATION_PATH,
f"nodes[{index}].out_map[{source_field!r}]", f"nodes[{index}].output[{output_index}].target",
"destination path must start with state.", "destination path must start with state. and reference a declared root field",
) )
if has_overlapping_paths(node.out_map.values()): if has_overlapping_paths(output_targets):
report.add( report.add(
ValidationIssueCode.INVALID_DESTINATION_PATH, ValidationIssueCode.INVALID_DESTINATION_PATH,
f"nodes[{index}].out_map", f"nodes[{index}].output",
"out_map has overlapping state destination paths", "output has overlapping state destination paths",
) )
def _local_root(path: str) -> str | None: def _local_root(path: str | LocalPath) -> str | None:
try: try:
parts = split_local_path(path) parts = split_local_path(path)
except LocalPathError: except LocalPathError:
@@ -113,6 +115,13 @@ def _local_root(path: str) -> str | None:
return "." if not parts else parts[0] return "." if not parts else parts[0]
def _state_destination_root(path: object) -> str | None:
try:
return StatePath.parse(str(path)).parts[0]
except PathResolutionError:
return None
def validate_condition_node( def validate_condition_node(
node: ConditionNode, node: ConditionNode,
index: int, index: int,
+185
View File
@@ -0,0 +1,185 @@
from __future__ import annotations
import pytest
from wf_core import (
END,
Edge,
NodeDef,
NodeUse,
ReducerRef,
SchemaRef,
StateField,
StateSchema,
Workflow,
WorkflowExecutionError,
)
from wf_core.models.steps import OutputBinding
from wf_core.runtime.engine import resume_workflow
from wf_core.runtime.ops.runs import create_run_state
from wf_core.runtime.ops.state import apply_output_bindings
def test_output_bindings_commit_patch_atomically_when_source_is_missing() -> None:
workflow = _workflow()
state = {"person": {"name": "old"}}
with pytest.raises(WorkflowExecutionError, match="missing"):
apply_output_bindings(
workflow,
[
_binding("person.name", "state.person.name"),
_binding("missing", "state.person.extra"),
],
{"person": {"name": "new"}},
state,
)
assert state["person"]["name"] == "old"
assert "extra" not in state["person"]
def test_output_bindings_reject_overlapping_write_targets_before_mutation() -> None:
workflow = _workflow()
state = {"person": {"name": "old"}}
with pytest.raises(WorkflowExecutionError, match="overlapping"):
apply_output_bindings(
workflow,
[
_binding("person", "state.person"),
_binding("person.name", "state.person.name"),
],
{"person": {"name": "Ada"}},
state,
)
assert state["person"]["name"] == "old"
def test_output_bindings_prepare_reducer_results_before_mutation() -> None:
workflow = _workflow(
fields={
"person.name": StateField(type="string"),
"person.tags": StateField(
type="array",
reducer=ReducerRef(name="wf.std.set_union", config={"bad": True}),
),
}
)
state = {"person": {"name": "old", "tags": ["seed"]}}
with pytest.raises(WorkflowExecutionError, match="reducer config"):
apply_output_bindings(
workflow,
[
_binding("person.name", "state.person.name"),
_binding("person.tags", "state.person.tags"),
],
{"person": {"name": "new", "tags": ["next"]}},
state,
)
assert state["person"]["name"] == "old"
assert state["person"]["tags"][0] == "seed"
assert len(state["person"]["tags"]) == 1
def test_output_bindings_commit_to_staged_state_before_mutating_original() -> None:
workflow = _workflow(
fields={
"person.name": StateField(type="string"),
"blocked.child": StateField(type="string"),
}
)
state = {"person": {"name": "old"}, "blocked": "not-an-object"}
with pytest.raises(WorkflowExecutionError, match="cannot descend"):
apply_output_bindings(
workflow,
[
_binding("person.name", "state.person.name"),
_binding("blocked.child", "state.blocked.child"),
],
{"person": {"name": "new"}, "blocked": {"child": "value"}},
state,
)
assert state["person"]["name"] == "old"
assert state["blocked"] == "not-an-object"
def test_full_workflow_execution_writes_canonical_output_bindings() -> None:
workflow = _workflow_with_node()
run = create_run_state(workflow, {})
run = resume_workflow(
workflow,
run,
{
"rename": lambda _payload, _ctx: {
"outcome": "ok",
"output": {"person": {"name": "Ada"}},
}
},
)
assert run.state["person"]["name"] == "Ada"
assert run.trace[0].state_changes["state.person.name"] == "Ada"
def _binding(source: str, target: str) -> OutputBinding:
return OutputBinding.model_validate({"source": source, "target": target})
def _workflow(
fields: dict[str, StateField] | None = None,
) -> Workflow:
return Workflow(
name="patch",
input_schema=SchemaRef(type="object", properties={}),
state_schema=StateSchema(
fields=fields
or {
"person": StateField(type="object"),
"person.name": StateField(type="string"),
"person.extra": StateField(type="string"),
}
),
output_schema=SchemaRef(type="object", properties={}),
start="n",
nodes=[],
edges=[],
)
def _workflow_with_node() -> Workflow:
return Workflow(
name="canonical_output",
input_schema=SchemaRef(type="object", properties={}),
state_schema=StateSchema(fields={"person.name": StateField(type="string")}),
output_schema=SchemaRef(type="object", properties={"person": {"type": "object"}}),
node_defs=[
NodeDef(
name="rename",
input_schema=SchemaRef(type="object", properties={}),
output_schema=SchemaRef(
type="object",
properties={"person": {"type": "object"}},
),
outcomes=["ok"],
)
],
start="rename",
nodes=[
NodeUse.model_validate(
{
"id": "rename",
"type": "node",
"node": "rename",
"output": [{"source": "person.name", "target": "state.person.name"}],
}
)
],
edges=[Edge.model_validate({"from": "rename", "outcome": "ok", "to": END})],
)
+166
View File
@@ -0,0 +1,166 @@
import pytest
from pydantic import ValidationError
from wf_core.models.steps import InputPathBinding, InputValueBinding, NodeUse
from wf_core.paths import GraphSourcePath, LocalPath, StatePath
def test_node_use_accepts_canonical_input_and_output_bindings():
node = NodeUse.model_validate(
{
"id": "echo",
"type": "node",
"node": "echo",
"input": [
{"target": "message", "path": "input.message"},
{"target": "mode", "value": None},
],
"output": [{"source": "echoed", "target": "state.echoed"}],
}
)
path_binding = node.input[0]
assert isinstance(path_binding, InputPathBinding)
assert path_binding.target == LocalPath.of("message")
assert path_binding.path == GraphSourcePath.input("message")
value_binding = node.input[1]
assert isinstance(value_binding, InputValueBinding)
assert value_binding.target == LocalPath.of("mode")
assert value_binding.value is None
assert node.output[0].source == LocalPath.of("echoed")
assert node.output[0].target == StatePath.of("echoed")
def test_node_use_converts_old_maps_to_canonical_bindings():
node = NodeUse.model_validate(
{
"id": "echo",
"type": "node",
"node": "echo",
"in_map": {"input.message": "message"},
"input_values": {"mode": "fast"},
"out_map": {"echoed": "state.echoed"},
}
)
dumped = node.model_dump(mode="json")
assert "in_map" not in dumped
assert "input_values" not in dumped
assert "out_map" not in dumped
assert dumped["input"][0]["value"] == "fast"
assert dumped["input"][0]["target"] == "mode"
assert dumped["input"][1]["path"] == "input.message"
assert dumped["input"][1]["target"] == "message"
assert dumped["output"][0]["source"] == "echoed"
assert dumped["output"][0]["target"] == "state.echoed"
def test_node_use_rejects_mixed_old_and_new_binding_styles():
with pytest.raises(ValidationError):
NodeUse.model_validate(
{
"id": "echo",
"type": "node",
"node": "echo",
"input": [{"target": "message", "path": "input.message"}],
"in_map": {"input.other": "other"},
}
)
def test_input_binding_rejects_path_and_value_together():
with pytest.raises(ValidationError):
NodeUse.model_validate(
{
"id": "bad",
"type": "node",
"node": "bad",
"input": [
{"target": "message", "path": "input.message", "value": "x"}
],
}
)
def test_input_binding_rejects_neither_path_nor_value():
with pytest.raises(ValidationError):
NodeUse.model_validate(
{
"id": "bad",
"type": "node",
"node": "bad",
"input": [{"target": "message"}],
}
)
@pytest.mark.parametrize(
"field,binding",
[
("input", {"target": "message", "path": "input.message", "extra": True}),
("output", {"source": "echoed", "target": "state.echoed", "extra": True}),
],
)
def test_bindings_reject_extra_fields(field: str, binding: dict[str, object]):
with pytest.raises(ValidationError):
NodeUse.model_validate(
{"id": "bad", "type": "node", "node": "bad", field: [binding]}
)
@pytest.mark.parametrize(
"field,value",
[
("in_map", None),
("input_values", []),
("out_map", "bad"),
],
)
def test_deprecated_maps_reject_non_mapping_values(field: str, value: object):
with pytest.raises(ValidationError):
NodeUse.model_validate(
{"id": "bad", "type": "node", "node": "bad", field: value}
)
def test_deprecated_conversion_preserves_input_value_then_in_map_order():
node = NodeUse.model_validate(
{
"id": "ordered",
"type": "node",
"node": "ordered",
"input_values": {"first": 1, "second": 2},
"in_map": {"input.third": "third", "state.fourth": "fourth"},
}
)
dumped_input = node.model_dump(mode="json")["input"]
assert dumped_input[0]["target"] == "first"
assert dumped_input[0]["value"] == 1
assert dumped_input[1]["target"] == "second"
assert dumped_input[1]["value"] == 2
assert dumped_input[2]["target"] == "third"
assert dumped_input[2]["path"] == "input.third"
assert dumped_input[3]["target"] == "fourth"
assert dumped_input[3]["path"] == "state.fourth"
def test_deprecated_input_value_preserves_explicit_null():
node = NodeUse.model_validate(
{
"id": "null",
"type": "node",
"node": "null",
"input_values": {"maybe": None},
}
)
value_binding = node.input[0]
assert isinstance(value_binding, InputValueBinding)
assert value_binding.value is None
dumped_input = node.model_dump(mode="json")["input"]
assert dumped_input[0]["target"] == "maybe"
assert dumped_input[0]["value"] is None
+139 -13
View File
@@ -1,6 +1,9 @@
from __future__ import annotations from __future__ import annotations
from typing import Any, cast
from wf_core import Edge, NodeDef, NodeUse, SchemaRef, StateField, StateSchema, Workflow from wf_core import Edge, NodeDef, NodeUse, SchemaRef, StateField, StateSchema, Workflow
from wf_core.validation.issues import ValidationIssueCode
def test_validation_allows_nested_node_local_paths() -> None: def test_validation_allows_nested_node_local_paths() -> None:
@@ -22,7 +25,10 @@ def test_validation_rejects_overlapping_node_input_destinations() -> None:
).validate_structure() ).validate_structure()
assert any( assert any(
"overlapping node-local input paths" in issue.message for issue in report.errors issue.code == ValidationIssueCode.INVALID_NODE_INPUT_FIELD
and issue.path == "nodes[0].input"
and "overlapping node-local input paths" in issue.message
for issue in report.errors
) )
@@ -36,18 +42,146 @@ def test_validation_rejects_overlapping_state_write_destinations() -> None:
).validate_structure() ).validate_structure()
assert any( assert any(
"overlapping state destination paths" in issue.message issue.code == ValidationIssueCode.INVALID_DESTINATION_PATH
and issue.path == "nodes[0].output"
and "overlapping state destination paths" in issue.message
for issue in report.errors for issue in report.errors
) )
def _workflow(*, in_map: dict[str, str], out_map: dict[str, str]) -> Workflow: def test_validation_rejects_invalid_canonical_input_source_path() -> None:
report = _workflow(
input=[{"target": "user.name", "path": "state.unknown.name"}],
output=[],
).validate_structure()
assert any(
issue.code == ValidationIssueCode.INVALID_SOURCE_PATH
and issue.path == "nodes[0].input[0].path"
for issue in report.errors
)
def test_validation_allows_canonical_input_source_under_declared_state_field_root() -> None:
report = _workflow(
input=[{"target": "user.name", "path": "state.person.name"}],
output=[],
state_fields={"person.name": StateField(type="string")},
).validate_structure()
assert not any(
issue.code == ValidationIssueCode.INVALID_SOURCE_PATH for issue in report.errors
)
def test_validation_rejects_invalid_canonical_output_destination() -> None:
workflow = _workflow(
input=[],
output=[{"source": "user.name", "target": "state.person.name"}],
)
# StatePath parsing rejects bad roots before workflow validation; mutate here so
# validate_node_use still guards malformed canonical destinations.
cast(Any, workflow.nodes[0]).output[0].target = "output.person.name"
report = workflow.validate_structure()
assert any(
issue.code == ValidationIssueCode.INVALID_DESTINATION_PATH
and issue.path == "nodes[0].output[0].target"
for issue in report.errors
)
def test_validation_rejects_undeclared_canonical_output_destination_root() -> None:
report = _workflow(
input=[],
output=[{"source": "user.name", "target": "state.unknown.foo"}],
).validate_structure()
assert any(
issue.code == ValidationIssueCode.INVALID_DESTINATION_PATH
and issue.path == "nodes[0].output[0].target"
for issue in report.errors
)
def test_validation_rejects_overlapping_canonical_input_targets() -> None:
report = _workflow(
input=[
{"target": "user", "value": {"name": "Ada"}},
{"target": "user.name", "path": "input.person.name"},
],
output=[],
).validate_structure()
assert any(
issue.code == ValidationIssueCode.INVALID_NODE_INPUT_FIELD
and issue.path == "nodes[0].input"
for issue in report.errors
)
def test_validation_rejects_overlapping_canonical_output_targets() -> None:
report = _workflow(
input=[],
output=[
{"source": "user", "target": "state.person"},
{"source": "user.name", "target": "state.person.name"},
],
).validate_structure()
assert any(
issue.code == ValidationIssueCode.INVALID_DESTINATION_PATH
and issue.path == "nodes[0].output"
for issue in report.errors
)
def test_validation_allows_valid_canonical_mapping() -> None:
report = _workflow(
input=[
{"target": "user.name", "path": "input.person.name"},
{"target": "user.nickname", "value": "Ada"},
],
output=[{"source": "user.age", "target": "state.person.age"}],
).validate_structure()
mapping_issue_codes = {
ValidationIssueCode.INVALID_NODE_INPUT_FIELD,
ValidationIssueCode.INVALID_NODE_OUTPUT_FIELD,
ValidationIssueCode.INVALID_SOURCE_PATH,
ValidationIssueCode.INVALID_DESTINATION_PATH,
}
assert not any(issue.code in mapping_issue_codes for issue in report.errors)
def _workflow(
*,
in_map: dict[str, str] | None = None,
out_map: dict[str, str] | None = None,
input: list[dict[str, object]] | None = None,
output: list[dict[str, str]] | None = None,
state_fields: dict[str, StateField] | None = None,
) -> Workflow:
node_data: dict[str, object] = {
"id": "tool",
"type": "node",
"node": "tool",
}
if input is not None or output is not None:
node_data["input"] = input or []
node_data["output"] = output or []
else:
node_data["in_map"] = in_map or {}
node_data["out_map"] = out_map or {}
return Workflow( return Workflow(
name="mapping_validation", name="mapping_validation",
input_schema=SchemaRef.model_validate( input_schema=SchemaRef.model_validate(
{"type": "object", "properties": {"person": {"type": "object"}}} {"type": "object", "properties": {"person": {"type": "object"}}}
), ),
state_schema=StateSchema(fields={"person": StateField(type="object")}), state_schema=StateSchema(
fields=state_fields or {"person": StateField(type="object")}
),
output_schema=SchemaRef(type="object", properties={}), output_schema=SchemaRef(type="object", properties={}),
node_defs=[ node_defs=[
NodeDef( NodeDef(
@@ -62,14 +196,6 @@ def _workflow(*, in_map: dict[str, str], out_map: dict[str, str]) -> Workflow:
) )
], ],
start="tool", start="tool",
nodes=[ nodes=[NodeUse.model_validate(node_data)],
NodeUse(
id="tool",
type="node",
node="tool",
in_map=in_map,
out_map=out_map,
)
],
edges=[Edge.model_validate({"from": "tool", "outcome": "ok", "to": "__end__"})], edges=[Edge.model_validate({"from": "tool", "outcome": "ok", "to": "__end__"})],
) )
+108 -24
View File
@@ -1,5 +1,7 @@
from __future__ import annotations from __future__ import annotations
from typing import Any, cast
import pytest import pytest
from wf_core import ( from wf_core import (
@@ -16,6 +18,72 @@ from wf_core import (
) )
def test_canonical_bindings_resolve_input_values_paths_and_explicit_null() -> None:
workflow = Workflow.model_validate(
{
"name": "canonical",
"input_schema": {
"type": "object",
"properties": {"message": {"type": "string"}},
},
"state_schema": {"fields": {"echoed": {"type": "string"}}},
"output_schema": {
"type": "object",
"properties": {"echoed": {"type": "string"}},
},
"start": "echo",
"node_defs": [
{
"name": "echo",
"input_schema": {
"type": "object",
"properties": {
"message": {"type": "string"},
"mode": {"type": "string"},
"maybe": {"type": "null"},
},
"required": ["message", "mode", "maybe"],
},
"output_schema": {
"type": "object",
"properties": {"echoed": {"type": "string"}},
},
"outcomes": ["ok"],
}
],
"nodes": [
{
"id": "echo",
"type": "node",
"node": "echo",
"input": [
{"target": "message", "path": "input.message"},
{"target": "mode", "value": "fast"},
{"target": "maybe", "value": None},
],
"output": [{"source": "echoed", "target": "state.echoed"}],
}
],
"edges": [{"from": "echo", "outcome": "ok", "to": END}],
}
)
run = execute_workflow(
workflow,
{"message": "hi"},
registry={
"echo": lambda payload, _ctx: {
"outcome": "ok",
"output": {"echoed": payload["message"]},
}
},
)
assert run.trace[0].resolved_input["message"] == "hi"
assert run.trace[0].resolved_input["mode"] == "fast"
assert run.trace[0].resolved_input["maybe"] is None
assert run.state["echoed"] == "hi"
def test_nested_node_local_paths_build_input_and_read_output() -> None: def test_nested_node_local_paths_build_input_and_read_output() -> None:
workflow = _nested_mapping_workflow() workflow = _nested_mapping_workflow()
@@ -33,9 +101,8 @@ def test_nested_node_local_paths_build_input_and_read_output() -> None:
}, },
) )
assert run.trace[0].resolved_input == { assert run.trace[0].resolved_input["user"]["name"] == "Ada"
"user": {"name": "Ada", "email": "[email protected]"} assert run.trace[0].resolved_input["user"]["email"] == "[email protected]"
}
assert run.state["person"]["age"] == 36 assert run.state["person"]["age"] == 36
assert run.state["person"]["gender"] == "x" assert run.state["person"]["gender"] == "x"
assert run.state["experience"]["years"] == 12 assert run.state["experience"]["years"] == 12
@@ -46,7 +113,7 @@ def test_missing_nested_node_output_path_fails() -> None:
with pytest.raises( with pytest.raises(
WorkflowExecutionError, WorkflowExecutionError,
match="did not return required mapped field 'user.gender'", match="node output did not include required field 'user.gender'",
): ):
execute_workflow( execute_workflow(
workflow, workflow,
@@ -87,12 +154,17 @@ def test_root_node_local_paths_map_whole_input_and_output_payloads() -> None:
], ],
start="force", start="force",
nodes=[ nodes=[
NodeUse( cast(
id="force", Any,
type="node", NodeUse.model_validate(
node="force_rates", {
in_map={"input.rates": "."}, "id": "force",
out_map={".": "state.rates"}, "type": "node",
"node": "force_rates",
"in_map": {"input.rates": "."},
"out_map": {".": "state.rates"},
}
),
) )
], ],
edges=[Edge.model_validate({"from": "force", "outcome": "ok", "to": END})], edges=[Edge.model_validate({"from": "force", "outcome": "ok", "to": END})],
@@ -109,8 +181,10 @@ def test_root_node_local_paths_map_whole_input_and_output_payloads() -> None:
}, },
) )
assert run.trace[0].resolved_input == {"r_1": 0.9, "r_10": 0.1} assert run.trace[0].resolved_input["r_1"] == 0.9
assert run.state["rates"] == {"r_1": 0.0, "r_10": 0.1} assert run.trace[0].resolved_input["r_10"] == 0.1
assert run.state["rates"]["r_1"] == 0.0
assert run.state["rates"]["r_10"] == 0.1
def test_static_input_values_are_merged_into_node_local_input() -> None: def test_static_input_values_are_merged_into_node_local_input() -> None:
@@ -141,12 +215,17 @@ def test_static_input_values_are_merged_into_node_local_input() -> None:
], ],
start="constant", start="constant",
nodes=[ nodes=[
NodeUse( cast(
id="constant", Any,
type="node", NodeUse.model_validate(
node="constant", {
input_values={"value": "CLICKED"}, "id": "constant",
out_map={"value": "state.message"}, "type": "node",
"node": "constant",
"input_values": {"value": "CLICKED"},
"out_map": {"value": "state.message"},
}
),
) )
], ],
edges=[Edge.model_validate({"from": "constant", "outcome": "ok", "to": END})], edges=[Edge.model_validate({"from": "constant", "outcome": "ok", "to": END})],
@@ -204,19 +283,24 @@ def _nested_mapping_workflow() -> Workflow:
], ],
start="big", start="big",
nodes=[ nodes=[
NodeUse( cast(
id="big", Any,
type="node", NodeUse.model_validate(
node="big_tool", {
in_map={ "id": "big",
"type": "node",
"node": "big_tool",
"in_map": {
"input.person.name": "user.name", "input.person.name": "user.name",
"input.digital.email": "user.email", "input.digital.email": "user.email",
}, },
out_map={ "out_map": {
"user.age": "state.person.age", "user.age": "state.person.age",
"user.gender": "state.person.gender", "user.gender": "state.person.gender",
"job.years": "state.experience.years", "job.years": "state.experience.years",
}, },
}
),
) )
], ],
edges=[Edge.model_validate({"from": "big", "outcome": "ok", "to": END})], edges=[Edge.model_validate({"from": "big", "outcome": "ok", "to": END})],
+185
View File
@@ -0,0 +1,185 @@
from __future__ import annotations
import pytest
from pydantic import BaseModel, ValidationError
from wf_core.paths import (
GraphSourcePath,
LocalPath,
PathResolutionError,
StatePath,
is_valid_destination_path,
is_valid_source_path,
)
def test_graph_source_path_accepts_root_and_nested_paths() -> None:
assert str(GraphSourcePath.parse("state")) == "state"
assert str(GraphSourcePath.parse("input")) == "input"
assert str(GraphSourcePath.parse("context")) == "context"
assert str(GraphSourcePath.parse("input.user")) == "input.user"
assert str(GraphSourcePath.parse("state.person.name")) == "state.person.name"
assert str(GraphSourcePath.context("loop_item")) == "context.loop_item"
def test_state_path_serializes_with_state_prefix() -> None:
assert str(StatePath.of("person.name")) == "state.person.name"
assert str(StatePath.parse("state.person.name")) == "state.person.name"
def test_state_path_rejects_bare_state_write_target() -> None:
with pytest.raises(PathResolutionError, match="state path"):
StatePath.parse("state")
def test_local_path_supports_root_marker_and_fragments() -> None:
assert str(LocalPath.root()) == "."
assert str(LocalPath.of("user.name")) == "user.name"
assert str(LocalPath.of("user", "name")) == "user.name"
assert LocalPath.parse(".") == LocalPath.root()
@pytest.mark.parametrize(
"raw",
[
"",
".",
"state.",
"state..name",
"state.items.0",
"state.user-name",
"state.items[0]",
"output.foo",
],
)
def test_graph_source_paths_reject_invalid_segments(raw: str) -> None:
with pytest.raises(PathResolutionError):
GraphSourcePath.parse(raw)
@pytest.mark.parametrize(
"factory",
[
LocalPath.parse,
StatePath.parse,
GraphSourcePath.parse,
],
)
@pytest.mark.parametrize(
"raw",
[
"state.",
"state..name",
"state.items.0",
"state.user-name",
"state.items[0]",
],
)
def test_all_path_types_reject_invalid_segments(factory, raw: str) -> None:
with pytest.raises(PathResolutionError):
factory(raw)
def test_path_objects_are_immutable_and_hashable() -> None:
paths = {StatePath.of("person.name"), StatePath.of("person.name")}
assert len(paths) == 1
with pytest.raises(Exception):
StatePath.of("person.name").parts = ("other",) # type: ignore[misc]
@pytest.mark.parametrize(
("factory", "args"),
[
(GraphSourcePath, ("output", ("user-name",))),
(StatePath, (("0",),)),
(LocalPath, (("items[0]",),)),
],
)
def test_direct_constructors_enforce_path_invariants(factory, args: tuple[object, ...]) -> None:
with pytest.raises(PathResolutionError):
factory(*args)
def test_pydantic_revalidates_existing_path_objects() -> None:
class Payload(BaseModel):
source: GraphSourcePath
target: StatePath
local: LocalPath
# Bypass constructors to simulate stale or malicious objects that predate
# constructor validation. Pydantic must not blindly trust existing instances.
source = object.__new__(GraphSourcePath)
object.__setattr__(source, "root", "output")
object.__setattr__(source, "parts", ("user-name",))
target = object.__new__(StatePath)
object.__setattr__(target, "parts", ("0",))
local = object.__new__(LocalPath)
object.__setattr__(local, "parts", ("items[0]",))
with pytest.raises(ValidationError):
Payload.model_validate({"source": source, "target": StatePath.of("person"), "local": LocalPath.root()})
with pytest.raises(ValidationError):
Payload.model_validate({"source": GraphSourcePath.input("user"), "target": target, "local": LocalPath.root()})
with pytest.raises(ValidationError):
Payload.model_validate({"source": GraphSourcePath.input("user"), "target": StatePath.of("person"), "local": local})
def test_pydantic_accepts_path_strings_and_serializes_strings() -> None:
class Payload(BaseModel):
source: GraphSourcePath
target: StatePath
local: LocalPath
payload = Payload.model_validate(
{"source": "input.user", "target": "state.person", "local": "user"}
)
assert payload.source == GraphSourcePath.input("user")
assert payload.target == StatePath.of("person")
assert payload.local == LocalPath.of("user")
dumped = payload.model_dump(mode="json")
assert dumped["source"] == "input.user"
assert dumped["target"] == "state.person"
assert dumped["local"] == "user"
def test_pydantic_accepts_existing_path_objects() -> None:
class Payload(BaseModel):
source: GraphSourcePath
target: StatePath
local: LocalPath
payload = Payload.model_validate(
{
"source": GraphSourcePath.state("person"),
"target": StatePath.of("person.name"),
"local": LocalPath.root(),
}
)
assert str(payload.source) == "state.person"
assert str(payload.target) == "state.person.name"
assert str(payload.local) == "."
def test_pydantic_rejects_bad_path_string() -> None:
class Payload(BaseModel):
source: GraphSourcePath
with pytest.raises(ValidationError):
Payload.model_validate({"source": "output.foo"})
def test_existing_source_and_destination_validation_helpers_use_new_parsers() -> None:
assert is_valid_source_path("state", set(), set()) is True
assert is_valid_source_path("input", set(), set()) is True
assert is_valid_source_path("context", set(), set(), allow_context=True) is True
assert is_valid_source_path("state.person", {"person"}, set()) is True
assert is_valid_source_path("input.person", set(), {"person"}) is True
assert is_valid_source_path("state.person-name", {"person-name"}, set()) is False
assert is_valid_destination_path("state") is False
assert is_valid_destination_path("state.person") is True
assert is_valid_destination_path("input.person") is False