interrupt
This commit is contained in:
@@ -3,13 +3,10 @@ import sys
|
|||||||
|
|
||||||
from wf_core import (
|
from wf_core import (
|
||||||
END,
|
END,
|
||||||
RunState,
|
|
||||||
RunStatus,
|
|
||||||
RuntimeContext,
|
RuntimeContext,
|
||||||
Workflow,
|
Workflow,
|
||||||
execute_workflow,
|
execute_workflow,
|
||||||
resume_workflow,
|
resume_workflow,
|
||||||
step_workflow,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
@@ -30,6 +27,8 @@ workflow = Workflow.model_validate(
|
|||||||
"should_email": {"type": "boolean"},
|
"should_email": {"type": "boolean"},
|
||||||
"documents": {"type": "array", "merge_strategy": "replace"},
|
"documents": {"type": "array", "merge_strategy": "replace"},
|
||||||
"summary": {"type": "string", "merge_strategy": "replace"},
|
"summary": {"type": "string", "merge_strategy": "replace"},
|
||||||
|
"approved": {"type": "boolean", "merge_strategy": "replace"},
|
||||||
|
"approval_comment": {"type": "string", "merge_strategy": "replace"},
|
||||||
"email_status": {"type": "string", "merge_strategy": "replace"},
|
"email_status": {"type": "string", "merge_strategy": "replace"},
|
||||||
}
|
}
|
||||||
},
|
},
|
||||||
@@ -133,6 +132,20 @@ workflow = Workflow.model_validate(
|
|||||||
"in_map": {"state.summary": "summary"},
|
"in_map": {"state.summary": "summary"},
|
||||||
"out_map": {"email_status": "state.email_status"},
|
"out_map": {"email_status": "state.email_status"},
|
||||||
},
|
},
|
||||||
|
{
|
||||||
|
"id": "approve_email",
|
||||||
|
"type": "interrupt",
|
||||||
|
"kind": "approval",
|
||||||
|
"request_map": {
|
||||||
|
"state.summary": "summary",
|
||||||
|
"input.folder_id": "folder_id",
|
||||||
|
},
|
||||||
|
"out_map": {
|
||||||
|
"approved": "state.approved",
|
||||||
|
"comment": "state.approval_comment",
|
||||||
|
},
|
||||||
|
"outcomes": ["submitted", "cancelled"],
|
||||||
|
},
|
||||||
{
|
{
|
||||||
"id": "skip_email",
|
"id": "skip_email",
|
||||||
"type": "node",
|
"type": "node",
|
||||||
@@ -144,8 +157,10 @@ workflow = Workflow.model_validate(
|
|||||||
"edges": [
|
"edges": [
|
||||||
{"from": "list_files", "outcome": "ok", "to": "summarize"},
|
{"from": "list_files", "outcome": "ok", "to": "summarize"},
|
||||||
{"from": "summarize", "outcome": "ok", "to": "should_email"},
|
{"from": "summarize", "outcome": "ok", "to": "should_email"},
|
||||||
{"from": "should_email", "outcome": "true", "to": "send_email"},
|
{"from": "should_email", "outcome": "true", "to": "approve_email"},
|
||||||
{"from": "should_email", "outcome": "false", "to": "skip_email"},
|
{"from": "should_email", "outcome": "false", "to": "skip_email"},
|
||||||
|
{"from": "approve_email", "outcome": "submitted", "to": "send_email"},
|
||||||
|
{"from": "approve_email", "outcome": "cancelled", "to": "skip_email"},
|
||||||
{"from": "send_email", "outcome": "sent", "to": END},
|
{"from": "send_email", "outcome": "sent", "to": END},
|
||||||
{"from": "skip_email", "outcome": "ok", "to": END},
|
{"from": "skip_email", "outcome": "ok", "to": END},
|
||||||
],
|
],
|
||||||
@@ -206,28 +221,30 @@ registry = {
|
|||||||
"mark_email_skipped": mark_email_skipped,
|
"mark_email_skipped": mark_email_skipped,
|
||||||
}
|
}
|
||||||
|
|
||||||
workflow_input = {"folder_id": "demo-folder", "should_email": False}
|
|
||||||
|
|
||||||
step_run = RunState(
|
|
||||||
workflow_name=workflow.name,
|
|
||||||
status=RunStatus.PENDING,
|
|
||||||
workflow_input=dict(workflow_input),
|
|
||||||
state=dict(workflow_input),
|
|
||||||
current_node_id=workflow.start,
|
|
||||||
)
|
|
||||||
|
|
||||||
workflow.validate_structure().raise_for_errors()
|
workflow.validate_structure().raise_for_errors()
|
||||||
|
|
||||||
print("Step-by-step run:")
|
print("Interrupting run:")
|
||||||
while step_run.current_node_id != END:
|
interrupted_run = execute_workflow(
|
||||||
step_workflow(workflow, step_run, registry)
|
workflow,
|
||||||
print(json.dumps(step_run.to_dict(), indent=2))
|
{"folder_id": "demo-folder", "should_email": True},
|
||||||
|
registry,
|
||||||
|
)
|
||||||
|
print(json.dumps(interrupted_run.to_dict(), indent=2))
|
||||||
|
|
||||||
step_run = resume_workflow(workflow, step_run, registry)
|
print("Resumed run:")
|
||||||
|
resumed_run = resume_workflow(
|
||||||
|
workflow,
|
||||||
|
interrupted_run,
|
||||||
|
registry,
|
||||||
|
resume_payload={"approved": True, "comment": "Looks good to send."},
|
||||||
|
resume_outcome="submitted",
|
||||||
|
)
|
||||||
|
print(json.dumps(resumed_run.to_dict(), indent=2))
|
||||||
|
|
||||||
print("Completed stepped run:")
|
print("Non-interrupt run:")
|
||||||
print(json.dumps(step_run.to_dict(), indent=2))
|
non_interrupt_run = execute_workflow(
|
||||||
|
workflow,
|
||||||
print("One-shot run:")
|
{"folder_id": "demo-folder", "should_email": False},
|
||||||
full_run = execute_workflow(workflow, workflow_input, registry)
|
registry,
|
||||||
print(json.dumps(full_run.to_dict(), indent=2))
|
)
|
||||||
|
print(json.dumps(non_interrupt_run.to_dict(), indent=2))
|
||||||
|
|||||||
+4
-1
@@ -2,6 +2,7 @@ from .model import (
|
|||||||
ConditionNode,
|
ConditionNode,
|
||||||
Edge,
|
Edge,
|
||||||
ForeachNode,
|
ForeachNode,
|
||||||
|
InterruptNode,
|
||||||
JoinNode,
|
JoinNode,
|
||||||
NodeDef,
|
NodeDef,
|
||||||
NodeResult,
|
NodeResult,
|
||||||
@@ -18,7 +19,7 @@ from .runtime import (
|
|||||||
resume_workflow,
|
resume_workflow,
|
||||||
step_workflow,
|
step_workflow,
|
||||||
)
|
)
|
||||||
from .run_state import RunState, RunStatus, RuntimeContext, TraceEntry
|
from .run_state import InterruptRequest, RunState, RunStatus, RuntimeContext, TraceEntry
|
||||||
from .tokens import END, START
|
from .tokens import END, START
|
||||||
from .validate import (
|
from .validate import (
|
||||||
ValidationIssue,
|
ValidationIssue,
|
||||||
@@ -31,6 +32,7 @@ __all__ = [
|
|||||||
"ConditionNode",
|
"ConditionNode",
|
||||||
"Edge",
|
"Edge",
|
||||||
"ForeachNode",
|
"ForeachNode",
|
||||||
|
"InterruptNode",
|
||||||
"JoinNode",
|
"JoinNode",
|
||||||
"NodeDef",
|
"NodeDef",
|
||||||
"NodeResult",
|
"NodeResult",
|
||||||
@@ -42,6 +44,7 @@ __all__ = [
|
|||||||
"RunStatus",
|
"RunStatus",
|
||||||
"RuntimeContext",
|
"RuntimeContext",
|
||||||
"TraceEntry",
|
"TraceEntry",
|
||||||
|
"InterruptRequest",
|
||||||
"START",
|
"START",
|
||||||
"END",
|
"END",
|
||||||
"ValidationIssue",
|
"ValidationIssue",
|
||||||
|
|||||||
+10
-1
@@ -104,8 +104,17 @@ class JoinNode(BaseModel):
|
|||||||
type: Literal["join"]
|
type: Literal["join"]
|
||||||
|
|
||||||
|
|
||||||
|
class InterruptNode(BaseModel):
|
||||||
|
id: str
|
||||||
|
type: Literal["interrupt"]
|
||||||
|
kind: str
|
||||||
|
request_map: dict[str, str] = Field(default_factory=dict)
|
||||||
|
out_map: dict[str, str] = Field(default_factory=dict)
|
||||||
|
outcomes: list[str] = Field(default_factory=lambda: ["submitted"])
|
||||||
|
|
||||||
|
|
||||||
Step = Annotated[
|
Step = Annotated[
|
||||||
NodeUse | ConditionNode | ForeachNode | JoinNode,
|
NodeUse | ConditionNode | ForeachNode | JoinNode | InterruptNode,
|
||||||
Field(discriminator="type"),
|
Field(discriminator="type"),
|
||||||
]
|
]
|
||||||
|
|
||||||
|
|||||||
@@ -32,6 +32,15 @@ class TraceEntry:
|
|||||||
state_changes: dict[str, Any] = field(default_factory=dict)
|
state_changes: dict[str, Any] = field(default_factory=dict)
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(slots=True)
|
||||||
|
class InterruptRequest:
|
||||||
|
id: str
|
||||||
|
node_id: str
|
||||||
|
kind: str
|
||||||
|
payload: dict[str, Any] = field(default_factory=dict)
|
||||||
|
resumable: bool = True
|
||||||
|
|
||||||
|
|
||||||
@dataclass(slots=True)
|
@dataclass(slots=True)
|
||||||
class RunState:
|
class RunState:
|
||||||
workflow_name: str
|
workflow_name: str
|
||||||
@@ -44,6 +53,7 @@ class RunState:
|
|||||||
prior_outcome: str | None = None
|
prior_outcome: str | None = None
|
||||||
activated_incoming_edge: str | None = None
|
activated_incoming_edge: str | None = None
|
||||||
error: str | None = None
|
error: str | None = None
|
||||||
|
interrupt: InterruptRequest | None = None
|
||||||
|
|
||||||
def to_dict(self) -> dict[str, Any]:
|
def to_dict(self) -> dict[str, Any]:
|
||||||
return asdict(self)
|
return asdict(self)
|
||||||
|
|||||||
+151
-8
@@ -5,10 +5,19 @@ from typing import Any
|
|||||||
|
|
||||||
from .conditions import eval_condition, safe_resolve_path
|
from .conditions import eval_condition, safe_resolve_path
|
||||||
from .errors import WorkflowExecutionError
|
from .errors import WorkflowExecutionError
|
||||||
from .model import ConditionNode, ForeachNode, JoinNode, NodeDef, NodeResult, NodeUse, Workflow
|
from .model import (
|
||||||
from .run_state import RunState, RunStatus, RuntimeContext, TraceEntry
|
ConditionNode,
|
||||||
|
ForeachNode,
|
||||||
|
InterruptNode,
|
||||||
|
JoinNode,
|
||||||
|
NodeDef,
|
||||||
|
NodeResult,
|
||||||
|
NodeUse,
|
||||||
|
Workflow,
|
||||||
|
)
|
||||||
|
from .run_state import InterruptRequest, RunState, RunStatus, RuntimeContext, TraceEntry
|
||||||
from .schema_tools import validate_payload_against_schema
|
from .schema_tools import validate_payload_against_schema
|
||||||
from .state_ops import apply_output_map, project_output
|
from .state_ops import apply_mapped_state, apply_output_map, project_output
|
||||||
from .tokens import END
|
from .tokens import END
|
||||||
|
|
||||||
|
|
||||||
@@ -44,6 +53,9 @@ def resume_workflow(
|
|||||||
workflow: Workflow,
|
workflow: Workflow,
|
||||||
run: RunState,
|
run: RunState,
|
||||||
registry: dict[str, NodeHandler],
|
registry: dict[str, NodeHandler],
|
||||||
|
*,
|
||||||
|
resume_payload: dict[str, Any] | None = None,
|
||||||
|
resume_outcome: str = "submitted",
|
||||||
) -> RunState:
|
) -> RunState:
|
||||||
if run.workflow_name != workflow.name:
|
if run.workflow_name != workflow.name:
|
||||||
raise WorkflowExecutionError(
|
raise WorkflowExecutionError(
|
||||||
@@ -56,14 +68,43 @@ def resume_workflow(
|
|||||||
if run.status == RunStatus.COMPLETED:
|
if run.status == RunStatus.COMPLETED:
|
||||||
return run
|
return run
|
||||||
|
|
||||||
run.status = RunStatus.RUNNING
|
|
||||||
run.error = None
|
|
||||||
node_defs = {node_def.name: node_def for node_def in workflow.node_defs}
|
node_defs = {node_def.name: node_def for node_def in workflow.node_defs}
|
||||||
nodes_by_id = {node.id: node for node in workflow.nodes}
|
nodes_by_id = {node.id: node for node in workflow.nodes}
|
||||||
edge_map = {(edge.from_, edge.outcome): edge.to for edge in workflow.edges}
|
edge_map = {(edge.from_, edge.outcome): edge.to for edge in workflow.edges}
|
||||||
|
|
||||||
|
if run.status == RunStatus.INTERRUPTED:
|
||||||
|
if resume_payload is None:
|
||||||
|
return run
|
||||||
|
_resume_interrupt(
|
||||||
|
workflow,
|
||||||
|
run,
|
||||||
|
nodes_by_id=nodes_by_id,
|
||||||
|
edge_map=edge_map,
|
||||||
|
resume_payload=resume_payload,
|
||||||
|
resume_outcome=resume_outcome,
|
||||||
|
)
|
||||||
|
if run.current_node_id == END:
|
||||||
|
run.output = project_output(workflow, run.state)
|
||||||
|
validate_payload_against_schema(
|
||||||
|
workflow.output_schema, run.output, "workflow output"
|
||||||
|
)
|
||||||
|
run.status = RunStatus.COMPLETED
|
||||||
|
return run
|
||||||
|
|
||||||
|
run.status = RunStatus.RUNNING
|
||||||
|
run.error = None
|
||||||
|
|
||||||
while run.current_node_id != END:
|
while run.current_node_id != END:
|
||||||
step_workflow(workflow, run, registry, node_defs=node_defs, nodes_by_id=nodes_by_id, edge_map=edge_map)
|
step_workflow(
|
||||||
|
workflow,
|
||||||
|
run,
|
||||||
|
registry,
|
||||||
|
node_defs=node_defs,
|
||||||
|
nodes_by_id=nodes_by_id,
|
||||||
|
edge_map=edge_map,
|
||||||
|
)
|
||||||
|
if run.status == RunStatus.INTERRUPTED:
|
||||||
|
return run
|
||||||
|
|
||||||
run.output = project_output(workflow, run.state)
|
run.output = project_output(workflow, run.state)
|
||||||
validate_payload_against_schema(
|
validate_payload_against_schema(
|
||||||
@@ -85,14 +126,20 @@ def step_workflow(
|
|||||||
) -> RunState:
|
) -> RunState:
|
||||||
if run.current_node_id is None or run.current_node_id == END:
|
if run.current_node_id is None or run.current_node_id == END:
|
||||||
return run
|
return run
|
||||||
|
if run.status == RunStatus.INTERRUPTED:
|
||||||
|
return run
|
||||||
|
|
||||||
if run.status == RunStatus.PENDING:
|
if run.status == RunStatus.PENDING:
|
||||||
run.status = RunStatus.RUNNING
|
run.status = RunStatus.RUNNING
|
||||||
run.error = None
|
run.error = None
|
||||||
|
|
||||||
node_defs = node_defs or {node_def.name: node_def for node_def in workflow.node_defs}
|
node_defs = node_defs or {
|
||||||
|
node_def.name: node_def for node_def in workflow.node_defs
|
||||||
|
}
|
||||||
nodes_by_id = nodes_by_id or {node.id: node for node in workflow.nodes}
|
nodes_by_id = nodes_by_id or {node.id: node for node in workflow.nodes}
|
||||||
edge_map = edge_map or {(edge.from_, edge.outcome): edge.to for edge in workflow.edges}
|
edge_map = edge_map or {
|
||||||
|
(edge.from_, edge.outcome): edge.to for edge in workflow.edges
|
||||||
|
}
|
||||||
|
|
||||||
step = nodes_by_id[run.current_node_id]
|
step = nodes_by_id[run.current_node_id]
|
||||||
|
|
||||||
@@ -117,6 +164,26 @@ def step_workflow(
|
|||||||
"output": {},
|
"output": {},
|
||||||
"state_changes": {},
|
"state_changes": {},
|
||||||
}
|
}
|
||||||
|
elif isinstance(step, InterruptNode):
|
||||||
|
interrupt_request = _build_interrupt_request(
|
||||||
|
step,
|
||||||
|
run.state,
|
||||||
|
run.workflow_input,
|
||||||
|
)
|
||||||
|
run.interrupt = interrupt_request
|
||||||
|
run.status = RunStatus.INTERRUPTED
|
||||||
|
run.trace.append(
|
||||||
|
TraceEntry(
|
||||||
|
node_id=run.current_node_id,
|
||||||
|
step_type=step.type,
|
||||||
|
resolved_input=interrupt_request.payload,
|
||||||
|
outcome="interrupt",
|
||||||
|
next_node_id=run.current_node_id,
|
||||||
|
output=interrupt_request.payload,
|
||||||
|
state_changes={},
|
||||||
|
)
|
||||||
|
)
|
||||||
|
return run
|
||||||
elif isinstance(step, ForeachNode):
|
elif isinstance(step, ForeachNode):
|
||||||
raise WorkflowExecutionError("foreach execution is not implemented yet")
|
raise WorkflowExecutionError("foreach execution is not implemented yet")
|
||||||
else:
|
else:
|
||||||
@@ -203,3 +270,79 @@ def coerce_node_result(raw_result: NodeResult | dict[str, Any]) -> NodeResult:
|
|||||||
if "outcome" in raw_result and "output" in raw_result:
|
if "outcome" in raw_result and "output" in raw_result:
|
||||||
return NodeResult.model_validate(raw_result)
|
return NodeResult.model_validate(raw_result)
|
||||||
return NodeResult(outcome="ok", output=raw_result)
|
return NodeResult(outcome="ok", output=raw_result)
|
||||||
|
|
||||||
|
|
||||||
|
def _build_interrupt_request(
|
||||||
|
node: InterruptNode,
|
||||||
|
state: dict[str, Any],
|
||||||
|
workflow_input: dict[str, Any],
|
||||||
|
) -> InterruptRequest:
|
||||||
|
payload = {
|
||||||
|
payload_field: safe_resolve_path(
|
||||||
|
source_path,
|
||||||
|
state=state,
|
||||||
|
workflow_input=workflow_input,
|
||||||
|
context={},
|
||||||
|
)
|
||||||
|
for source_path, payload_field in node.request_map.items()
|
||||||
|
}
|
||||||
|
return InterruptRequest(
|
||||||
|
id=f"interrupt:{node.id}",
|
||||||
|
node_id=node.id,
|
||||||
|
kind=node.kind,
|
||||||
|
payload=payload,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _resume_interrupt(
|
||||||
|
workflow: Workflow,
|
||||||
|
run: RunState,
|
||||||
|
*,
|
||||||
|
nodes_by_id: dict[str, Any],
|
||||||
|
edge_map: dict[tuple[str, str], str],
|
||||||
|
resume_payload: dict[str, Any],
|
||||||
|
resume_outcome: str,
|
||||||
|
) -> None:
|
||||||
|
if run.current_node_id is None:
|
||||||
|
raise WorkflowExecutionError("interrupted run has no current node")
|
||||||
|
if run.interrupt is None:
|
||||||
|
raise WorkflowExecutionError("run is interrupted but has no interrupt request")
|
||||||
|
|
||||||
|
step = nodes_by_id[run.current_node_id]
|
||||||
|
if not isinstance(step, InterruptNode):
|
||||||
|
raise WorkflowExecutionError(
|
||||||
|
f"interrupted run expected interrupt node, got {step.type!r}"
|
||||||
|
)
|
||||||
|
if resume_outcome not in step.outcomes:
|
||||||
|
raise WorkflowExecutionError(
|
||||||
|
f"interrupt node {step.id!r} does not declare resume outcome {resume_outcome!r}"
|
||||||
|
)
|
||||||
|
|
||||||
|
state_changes = apply_mapped_state(
|
||||||
|
workflow,
|
||||||
|
resume_payload,
|
||||||
|
step.out_map,
|
||||||
|
run.state,
|
||||||
|
missing_field_message="interrupt resume payload is missing required field {field}",
|
||||||
|
)
|
||||||
|
next_node_id = edge_map.get((run.current_node_id, resume_outcome))
|
||||||
|
if next_node_id is None:
|
||||||
|
raise WorkflowExecutionError(
|
||||||
|
f"no edge found for interrupt node {run.current_node_id!r} and outcome {resume_outcome!r}"
|
||||||
|
)
|
||||||
|
|
||||||
|
run.trace.append(
|
||||||
|
TraceEntry(
|
||||||
|
node_id=run.current_node_id,
|
||||||
|
step_type=step.type,
|
||||||
|
resolved_input=resume_payload,
|
||||||
|
outcome=resume_outcome,
|
||||||
|
next_node_id=next_node_id,
|
||||||
|
output=resume_payload,
|
||||||
|
state_changes=state_changes,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
run.prior_outcome = resume_outcome
|
||||||
|
run.activated_incoming_edge = run.current_node_id
|
||||||
|
run.current_node_id = next_node_id
|
||||||
|
run.interrupt = None
|
||||||
|
|||||||
+29
-7
@@ -4,7 +4,12 @@ from typing import Any
|
|||||||
|
|
||||||
from .errors import WorkflowExecutionError
|
from .errors import WorkflowExecutionError
|
||||||
from .model import NodeUse, Workflow
|
from .model import NodeUse, Workflow
|
||||||
from .paths import PathResolutionError, get_nested_value, set_nested_value, split_graph_path
|
from .paths import (
|
||||||
|
PathResolutionError,
|
||||||
|
get_nested_value,
|
||||||
|
set_nested_value,
|
||||||
|
split_graph_path,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def apply_output_map(
|
def apply_output_map(
|
||||||
@@ -13,13 +18,30 @@ def apply_output_map(
|
|||||||
node_output: dict[str, Any],
|
node_output: dict[str, Any],
|
||||||
state: dict[str, Any],
|
state: dict[str, Any],
|
||||||
) -> dict[str, Any]:
|
) -> dict[str, Any]:
|
||||||
state_changes: dict[str, Any] = {}
|
return apply_mapped_state(
|
||||||
for source_field, destination_path in node.out_map.items():
|
workflow,
|
||||||
if source_field not in node_output:
|
node_output,
|
||||||
raise WorkflowExecutionError(
|
node.out_map,
|
||||||
f"node {node.id!r} did not return required mapped field {source_field!r}"
|
state,
|
||||||
|
missing_field_message=f"node {node.id!r} did not return required mapped field {{field}}",
|
||||||
)
|
)
|
||||||
value = node_output[source_field]
|
|
||||||
|
|
||||||
|
def apply_mapped_state(
|
||||||
|
workflow: Workflow,
|
||||||
|
source_data: dict[str, Any],
|
||||||
|
mapping: dict[str, str],
|
||||||
|
state: dict[str, Any],
|
||||||
|
*,
|
||||||
|
missing_field_message: str,
|
||||||
|
) -> dict[str, Any]:
|
||||||
|
state_changes: dict[str, Any] = {}
|
||||||
|
for source_field, destination_path in mapping.items():
|
||||||
|
if source_field not in source_data:
|
||||||
|
raise WorkflowExecutionError(
|
||||||
|
missing_field_message.format(field=repr(source_field))
|
||||||
|
)
|
||||||
|
value = source_data[source_field]
|
||||||
write_state_value(workflow, state, destination_path, value)
|
write_state_value(workflow, state, destination_path, value)
|
||||||
state_changes[destination_path] = value
|
state_changes[destination_path] = value
|
||||||
return state_changes
|
return state_changes
|
||||||
|
|||||||
+45
-1
@@ -2,7 +2,6 @@ from __future__ import annotations
|
|||||||
|
|
||||||
from dataclasses import dataclass, field
|
from dataclasses import dataclass, field
|
||||||
from enum import StrEnum
|
from enum import StrEnum
|
||||||
from typing import Any
|
|
||||||
|
|
||||||
from .model import (
|
from .model import (
|
||||||
BinaryCondition,
|
BinaryCondition,
|
||||||
@@ -11,6 +10,7 @@ from .model import (
|
|||||||
Edge,
|
Edge,
|
||||||
ExistsCondition,
|
ExistsCondition,
|
||||||
ForeachNode,
|
ForeachNode,
|
||||||
|
InterruptNode,
|
||||||
LiteralOperand,
|
LiteralOperand,
|
||||||
NodeDef,
|
NodeDef,
|
||||||
NodeUse,
|
NodeUse,
|
||||||
@@ -41,6 +41,8 @@ class ValidationIssueCode(StrEnum):
|
|||||||
EMPTY_CONDITION_ARGS = "empty_condition_args"
|
EMPTY_CONDITION_ARGS = "empty_condition_args"
|
||||||
INVALID_CONDITION_PATH = "invalid_condition_path"
|
INVALID_CONDITION_PATH = "invalid_condition_path"
|
||||||
INVALID_FOREACH_SOURCE = "invalid_foreach_source"
|
INVALID_FOREACH_SOURCE = "invalid_foreach_source"
|
||||||
|
INVALID_INTERRUPT_SOURCE = "invalid_interrupt_source"
|
||||||
|
INVALID_INTERRUPT_DESTINATION = "invalid_interrupt_destination"
|
||||||
|
|
||||||
|
|
||||||
@dataclass(slots=True)
|
@dataclass(slots=True)
|
||||||
@@ -108,6 +110,10 @@ def validate_workflow(workflow: Workflow) -> ValidationReport:
|
|||||||
_validate_foreach_node(
|
_validate_foreach_node(
|
||||||
node, index, report, state_root_fields, input_root_fields
|
node, index, report, state_root_fields, input_root_fields
|
||||||
)
|
)
|
||||||
|
elif isinstance(node, InterruptNode):
|
||||||
|
_validate_interrupt_node(
|
||||||
|
node, index, report, state_root_fields, input_root_fields
|
||||||
|
)
|
||||||
|
|
||||||
if workflow.start not in nodes_by_id:
|
if workflow.start not in nodes_by_id:
|
||||||
report.add(
|
report.add(
|
||||||
@@ -258,6 +264,42 @@ def _validate_foreach_node(
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _validate_interrupt_node(
|
||||||
|
node: InterruptNode,
|
||||||
|
index: int,
|
||||||
|
report: ValidationReport,
|
||||||
|
state_root_fields: set[str],
|
||||||
|
input_root_fields: set[str],
|
||||||
|
) -> None:
|
||||||
|
for source_path, payload_field in node.request_map.items():
|
||||||
|
if not payload_field:
|
||||||
|
report.add(
|
||||||
|
ValidationIssueCode.INVALID_INTERRUPT_SOURCE,
|
||||||
|
f"nodes[{index}].request_map[{source_path!r}]",
|
||||||
|
"interrupt request payload field must not be empty",
|
||||||
|
)
|
||||||
|
if not is_valid_source_path(source_path, state_root_fields, input_root_fields):
|
||||||
|
report.add(
|
||||||
|
ValidationIssueCode.INVALID_INTERRUPT_SOURCE,
|
||||||
|
f"nodes[{index}].request_map[{source_path!r}]",
|
||||||
|
"interrupt request source must start with input. or state. and reference a declared root field",
|
||||||
|
)
|
||||||
|
|
||||||
|
for resume_field, destination_path in node.out_map.items():
|
||||||
|
if not resume_field:
|
||||||
|
report.add(
|
||||||
|
ValidationIssueCode.INVALID_INTERRUPT_DESTINATION,
|
||||||
|
f"nodes[{index}].out_map[{resume_field!r}]",
|
||||||
|
"interrupt resume field must not be empty",
|
||||||
|
)
|
||||||
|
if not is_valid_destination_path(destination_path):
|
||||||
|
report.add(
|
||||||
|
ValidationIssueCode.INVALID_INTERRUPT_DESTINATION,
|
||||||
|
f"nodes[{index}].out_map[{resume_field!r}]",
|
||||||
|
"interrupt resume destination must start with state.",
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def _validate_condition_expr(
|
def _validate_condition_expr(
|
||||||
condition: Condition,
|
condition: Condition,
|
||||||
path: str,
|
path: str,
|
||||||
@@ -346,6 +388,8 @@ def _declared_outcomes_for_step(step: Step, node_defs: dict[str, NodeDef]) -> se
|
|||||||
return {"done"}
|
return {"done"}
|
||||||
if step.type == "join":
|
if step.type == "join":
|
||||||
return {"done"}
|
return {"done"}
|
||||||
|
if step.type == "interrupt":
|
||||||
|
return set(step.outcomes)
|
||||||
return set()
|
return set()
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user