From a6b49de3f09153edacbdbe629643fb5ed59555e4 Mon Sep 17 00:00:00 2001 From: -LAN- Date: Fri, 7 Aug 2026 14:53:10 +0800 Subject: [PATCH] refactor(engine)!: simplify execution architecture Consolidate frame scheduling, event processing, commands, layers, filters, workers, and runtime execution state behind direct module boundaries and canonical public imports. BREAKING CHANGE: Rename the engine, runtime, event, layer, filter, command, and container APIs; remove legacy import paths and compatibility aliases; unify nested event ownership under container_id. --- README.md | 22 +- examples/slim_llm/code.py | 23 +- examples/slim_llm/dsl.py | 14 +- src/graphon/dsl/importer.py | 28 +- src/graphon/dsl/node_factory.py | 6 +- src/graphon/engine/__init__.py | 3 + .../command}/README.md | 16 +- src/graphon/engine/command/__init__.py | 23 + .../engine/command/builtin/__init__.py | 1 + .../command/builtin/in_memory.py} | 12 +- .../command/builtin/redis.py} | 116 ++- src/graphon/engine/command/entities.py | 42 + src/graphon/engine/command/processor.py | 109 +++ src/graphon/engine/command/protocol.py | 39 + .../engine/container_handler/__init__.py | 13 + .../container_handler/builtin/__init__.py | 1 + .../container_handler/builtin/iteration.py} | 100 ++- .../container_handler/builtin/loop.py} | 139 ++-- .../container_handler_factory.py | 8 + .../engine/container_handler/protocol.py | 62 ++ .../orchestration => engine}/dispatcher.py | 95 ++- .../graph_engine.py => engine/engine.py} | 235 +++--- src/graphon/engine/event/__init__.py | 1 + .../event/node_failure.py} | 102 +-- .../event/processor.py} | 138 ++-- .../event/stream.py} | 83 +- src/graphon/engine/filter/__init__.py | 13 + src/graphon/engine/filter/builtin/__init__.py | 1 + .../filter/builtin}/response_stream.py | 218 +++--- .../filters => engine/filter}/chain.py | 28 +- src/graphon/engine/filter/protocol.py | 40 + src/graphon/engine/frame.py | 278 +++++++ src/graphon/engine/layer/README.md | 40 + src/graphon/engine/layer/__init__.py | 13 + .../layers => engine/layer}/base.py | 54 +- src/graphon/engine/layer/builtin/__init__.py | 1 + .../layer/builtin}/execution_limits.py | 77 +- .../ready_queue/__init__.py | 6 +- .../ready_queue/entities.py} | 25 +- .../ready_queue/in_memory.py | 33 +- src/graphon/engine/scheduler.py | 402 ++++++++++ src/graphon/engine/worker/__init__.py | 12 + .../{graph_engine => engine/worker}/worker.py | 117 ++- src/graphon/engine/worker/worker_pool.py | 123 +++ .../__init__.py | 11 +- .../{graph_events => engine_events}/agent.py | 4 +- .../{graph_events => engine_events}/base.py | 21 +- .../{graph_events => engine_events}/graph.py | 14 +- .../iteration.py | 10 +- .../{graph_events => engine_events}/loop.py | 10 +- .../{graph_events => engine_events}/node.py | 28 +- .../traversal.py | 5 +- src/graphon/entities/__init__.py | 4 +- src/graphon/entities/graph_init_params.py | 4 +- src/graphon/entities/workflow_execution.py | 2 + src/graphon/filters/__init__.py | 15 - src/graphon/graph/graph.py | 180 ++++- src/graphon/graph/graph_template.py | 2 +- src/graphon/graph_engine/__init__.py | 22 - src/graphon/graph_engine/_engine_utils.py | 20 - .../graph_engine/command_channels/__init__.py | 7 - .../graph_engine/command_channels/protocol.py | 42 - .../command_processing/__init__.py | 19 - .../command_processing/command_handlers.py | 71 -- .../command_processing/command_processor.py | 110 --- src/graphon/graph_engine/config.py | 14 - .../graph_engine/container_handlers.py | 57 -- src/graphon/graph_engine/domain/__init__.py | 13 - .../graph_engine/domain/node_execution.py | 19 - src/graphon/graph_engine/entities/commands.py | 70 -- src/graphon/graph_engine/entities/tasks.py | 21 - .../graph_engine/event_management/__init__.py | 13 - src/graphon/graph_engine/filters/__init__.py | 15 - src/graphon/graph_engine/filters/base.py | 70 -- src/graphon/graph_engine/frames.py | 146 ---- .../graph_engine/graph_state_manager.py | 214 ------ .../graph_engine/graph_traversal/__init__.py | 13 - .../graph_traversal/edge_processor.py | 174 ----- .../graph_traversal/skip_propagator.py | 128 ---- src/graphon/graph_engine/layers/README.md | 52 -- src/graphon/graph_engine/layers/__init__.py | 15 - .../graph_engine/layers/debug_logging.py | 300 -------- src/graphon/graph_engine/manager.py | 85 -- .../graph_engine/orchestration/__init__.py | 9 - .../worker_management/__init__.py | 11 - .../worker_management/worker_pool.py | 278 ------- src/graphon/graph_events/human_input.py | 0 src/graphon/node_events/__init__.py | 4 +- src/graphon/node_events/agent.py | 4 +- src/graphon/node_events/base.py | 4 +- src/graphon/node_events/iteration.py | 10 +- src/graphon/node_events/loop.py | 10 +- src/graphon/node_events/node.py | 24 +- src/graphon/nodes/base/node.py | 85 +- src/graphon/nodes/code/code_node.py | 8 +- src/graphon/nodes/container_effects.py | 11 +- src/graphon/nodes/document_extractor/node.py | 8 +- src/graphon/nodes/http_request/node.py | 8 +- .../nodes/human_input/human_input_node.py | 12 +- src/graphon/nodes/iteration/iteration_node.py | 10 +- src/graphon/nodes/llm/node.py | 30 +- src/graphon/nodes/loop/loop_node.py | 8 +- .../parameter_extractor_node.py | 16 +- .../question_classifier_node.py | 8 +- .../template_transform_node.py | 8 +- src/graphon/nodes/tool/tool_node.py | 40 +- .../nodes/variable_assigner/v1/node.py | 12 +- .../nodes/variable_assigner/v2/node.py | 12 +- src/graphon/runtime/__init__.py | 4 +- src/graphon/runtime/container_state.py | 37 +- .../execution.py} | 72 +- src/graphon/runtime/graph_runtime_state.py | 206 +++-- .../runtime/graph_runtime_state_protocol.py | 2 +- src/graphon/runtime/read_only_wrappers.py | 6 +- src/graphon/runtime/ready_queue.py | 2 +- tests/dsl/test_node_factory.py | 19 +- .../entities => tests/engine}/__init__.py | 0 .../test_cooperative_container_execution.py | 105 ++- .../test_dispatch_patterns.py | 724 +++++++++++------- .../test_event_filters.py | 61 +- .../test_raw_engine_events.py | 108 ++- .../test_response_stream_filter.py | 30 +- .../test_scheduler.py} | 74 +- .../test_serializable_graph_runtime.py | 178 +++-- .../__init__.py | 0 .../test_traversal_events.py | 4 +- tests/graph/test_graph_scoping.py | 245 ++++++ tests/graph/test_graph_validation.py | 19 +- tests/graph_engine/__init__.py | 0 .../graph_engine/graph_traversal/__init__.py | 0 .../graph_traversal/test_skip_propagator.py | 225 ------ tests/helpers/builders.py | 6 +- tests/helpers/workflow_events.py | 26 +- tests/http/test_client.py | 7 +- tests/node_events/test_node_event_aliases.py | 6 +- .../nodes/base/test_node_execution_binding.py | 5 +- .../human_input/test_human_input_node.py | 7 +- tests/nodes/if_else/test_if_else_node.py | 7 +- tests/nodes/llm/test_node.py | 25 +- .../nodes/parameter_extractor/test_prompts.py | 14 +- .../test_question_classifier_node.py | 11 +- .../nodes/test_human_input_runtime_binding.py | 5 +- tests/nodes/tool/test_tool_node.py | 17 +- tests/nodes/variable_assigner/test_v1_node.py | 10 +- tests/nodes/variable_assigner/test_v2_node.py | 9 +- tests/runtime/test_custom_container_state.py | 83 ++ tests/runtime/test_graph_runtime_state.py | 164 +++- tests/test_protocol_abstract_contracts.py | 132 ---- ...h_events.py => test_full_engine_events.py} | 138 ++-- 149 files changed, 3991 insertions(+), 4364 deletions(-) create mode 100644 src/graphon/engine/__init__.py rename src/graphon/{graph_engine/command_channels => engine/command}/README.md (52%) create mode 100644 src/graphon/engine/command/__init__.py create mode 100644 src/graphon/engine/command/builtin/__init__.py rename src/graphon/{graph_engine/command_channels/in_memory_channel.py => engine/command/builtin/in_memory.py} (76%) rename src/graphon/{graph_engine/command_channels/redis_channel.py => engine/command/builtin/redis.py} (58%) create mode 100644 src/graphon/engine/command/entities.py create mode 100644 src/graphon/engine/command/processor.py create mode 100644 src/graphon/engine/command/protocol.py create mode 100644 src/graphon/engine/container_handler/__init__.py create mode 100644 src/graphon/engine/container_handler/builtin/__init__.py rename src/graphon/{graph_engine/iteration_container_handler.py => engine/container_handler/builtin/iteration.py} (85%) rename src/graphon/{graph_engine/loop_container_handler.py => engine/container_handler/builtin/loop.py} (78%) create mode 100644 src/graphon/engine/container_handler/container_handler_factory.py create mode 100644 src/graphon/engine/container_handler/protocol.py rename src/graphon/{graph_engine/orchestration => engine}/dispatcher.py (58%) rename src/graphon/{graph_engine/graph_engine.py => engine/engine.py} (58%) create mode 100644 src/graphon/engine/event/__init__.py rename src/graphon/{graph_engine/error_handler.py => engine/event/node_failure.py} (66%) rename src/graphon/{graph_engine/event_management/event_handlers.py => engine/event/processor.py} (74%) rename src/graphon/{graph_engine/event_management/event_manager.py => engine/event/stream.py} (69%) create mode 100644 src/graphon/engine/filter/__init__.py create mode 100644 src/graphon/engine/filter/builtin/__init__.py rename src/graphon/{graph_engine/filters => engine/filter/builtin}/response_stream.py (81%) rename src/graphon/{graph_engine/filters => engine/filter}/chain.py (63%) create mode 100644 src/graphon/engine/filter/protocol.py create mode 100644 src/graphon/engine/frame.py create mode 100644 src/graphon/engine/layer/README.md create mode 100644 src/graphon/engine/layer/__init__.py rename src/graphon/{graph_engine/layers => engine/layer}/base.py (70%) create mode 100644 src/graphon/engine/layer/builtin/__init__.py rename src/graphon/{graph_engine/layers => engine/layer/builtin}/execution_limits.py (64%) rename src/graphon/{graph_engine => engine}/ready_queue/__init__.py (64%) rename src/graphon/{graph_engine/ready_queue/protocol.py => engine/ready_queue/entities.py} (50%) rename src/graphon/{graph_engine => engine}/ready_queue/in_memory.py (85%) create mode 100644 src/graphon/engine/scheduler.py create mode 100644 src/graphon/engine/worker/__init__.py rename src/graphon/{graph_engine => engine/worker}/worker.py (77%) create mode 100644 src/graphon/engine/worker/worker_pool.py rename src/graphon/{graph_events => engine_events}/__init__.py (93%) rename src/graphon/{graph_events => engine_events}/agent.py (85%) rename src/graphon/{graph_events => engine_events}/base.py (52%) rename src/graphon/{graph_events => engine_events}/graph.py (83%) rename src/graphon/{graph_events => engine_events}/iteration.py (80%) rename src/graphon/{graph_events => engine_events}/loop.py (82%) rename src/graphon/{graph_events => engine_events}/node.py (86%) rename src/graphon/{graph_events => engine_events}/traversal.py (81%) delete mode 100644 src/graphon/filters/__init__.py delete mode 100644 src/graphon/graph_engine/__init__.py delete mode 100644 src/graphon/graph_engine/_engine_utils.py delete mode 100644 src/graphon/graph_engine/command_channels/__init__.py delete mode 100644 src/graphon/graph_engine/command_channels/protocol.py delete mode 100644 src/graphon/graph_engine/command_processing/__init__.py delete mode 100644 src/graphon/graph_engine/command_processing/command_handlers.py delete mode 100644 src/graphon/graph_engine/command_processing/command_processor.py delete mode 100644 src/graphon/graph_engine/config.py delete mode 100644 src/graphon/graph_engine/container_handlers.py delete mode 100644 src/graphon/graph_engine/domain/__init__.py delete mode 100644 src/graphon/graph_engine/domain/node_execution.py delete mode 100644 src/graphon/graph_engine/entities/commands.py delete mode 100644 src/graphon/graph_engine/entities/tasks.py delete mode 100644 src/graphon/graph_engine/event_management/__init__.py delete mode 100644 src/graphon/graph_engine/filters/__init__.py delete mode 100644 src/graphon/graph_engine/filters/base.py delete mode 100644 src/graphon/graph_engine/frames.py delete mode 100644 src/graphon/graph_engine/graph_state_manager.py delete mode 100644 src/graphon/graph_engine/graph_traversal/__init__.py delete mode 100644 src/graphon/graph_engine/graph_traversal/edge_processor.py delete mode 100644 src/graphon/graph_engine/graph_traversal/skip_propagator.py delete mode 100644 src/graphon/graph_engine/layers/README.md delete mode 100644 src/graphon/graph_engine/layers/__init__.py delete mode 100644 src/graphon/graph_engine/layers/debug_logging.py delete mode 100644 src/graphon/graph_engine/manager.py delete mode 100644 src/graphon/graph_engine/orchestration/__init__.py delete mode 100644 src/graphon/graph_engine/worker_management/__init__.py delete mode 100644 src/graphon/graph_engine/worker_management/worker_pool.py delete mode 100644 src/graphon/graph_events/human_input.py rename src/graphon/{graph_engine/domain/graph_execution.py => runtime/execution.py} (79%) rename {src/graphon/graph_engine/entities => tests/engine}/__init__.py (100%) rename tests/{graph_engine => engine}/test_cooperative_container_execution.py (82%) rename tests/{graph_engine => engine}/test_dispatch_patterns.py (70%) rename tests/{graph_engine => engine}/test_event_filters.py (58%) rename tests/{graph_engine => engine}/test_raw_engine_events.py (51%) rename tests/{graph_engine => engine}/test_response_stream_filter.py (97%) rename tests/{graph_engine/graph_traversal/test_traversal_events.py => engine/test_scheduler.py} (54%) rename tests/{graph_engine => engine}/test_serializable_graph_runtime.py (85%) rename tests/{graph_events => engine_events}/__init__.py (100%) rename tests/{graph_events => engine_events}/test_traversal_events.py (85%) create mode 100644 tests/graph/test_graph_scoping.py delete mode 100644 tests/graph_engine/__init__.py delete mode 100644 tests/graph_engine/graph_traversal/__init__.py delete mode 100644 tests/graph_engine/graph_traversal/test_skip_propagator.py create mode 100644 tests/runtime/test_custom_container_state.py delete mode 100644 tests/test_protocol_abstract_contracts.py rename tests/workflows/{test_full_graph_events.py => test_full_engine_events.py} (89%) diff --git a/README.md b/README.md index 0b6a8498..84688b2e 100644 --- a/README.md +++ b/README.md @@ -8,9 +8,9 @@ protocols, and a runnable end-to-end example. ## Highlights -- Queue-based `GraphEngine` orchestration with event-driven execution +- Queue-based `Engine` orchestration with event-driven execution - Graph parsing, validation, and fluent graph building -- Shared runtime state, variable pool, and workflow execution domain models +- Shared runtime state, variable pool, and workflow execution state - Built-in node implementations for common workflow patterns - DSL import support with Slim-backed LLM nodes - HTTP, file, tool, and human-input integration protocols @@ -89,9 +89,10 @@ For the exact credential shape and runtime notes, see At a high level, direct Graphon usage looks like this: 1. Build or load a graph and instantiate nodes into a `Graph`. -2. Prepare `GraphRuntimeState` and seed the `VariablePool`. +2. Prepare `RuntimeState` with the workflow ID and seed the `VariablePool`. 3. Configure model, file, HTTP, tool, or human-input adapters as needed. -4. Run `GraphEngine` and consume emitted graph events. +4. Run `Engine` and consume emitted engine events; a local command channel is + created automatically unless an external channel is supplied. 5. Read final outputs from runtime state. For Dify DSL documents, use `graphon.dsl.loads()` to build the engine from the @@ -128,14 +129,15 @@ planned as a separate follow-up. ## Project Layout - `src/graphon/graph`: graph structures, parsing, validation, and builders -- `src/graphon/graph_engine`: orchestration, workers, command channels, and - layers +- `src/graphon/engine`: dispatch, workers, commands, events, and layers - `src/graphon/runtime`: runtime state, read-only wrappers, and variable pool - `src/graphon/nodes`: built-in workflow node implementations - `src/graphon/model_runtime`: provider/model abstractions and shared model entities - `src/graphon/dsl`: DSL import support, including Slim-backed runtime adapters -- `src/graphon/graph_events`: event models emitted during execution +- `src/graphon/node_events`: payloads emitted by node implementations before + execution context is attached +- `src/graphon/engine_events`: complete events emitted by the engine - `src/graphon/http`: HTTP client abstractions and default implementation - `src/graphon/file`: workflow file models and file runtime helpers - `src/graphon/protocols`: public protocol re-exports for integrations @@ -149,10 +151,10 @@ planned as a separate follow-up. runnable Slim LLM example setup - [src/graphon/model_runtime/README.md](src/graphon/model_runtime/README.md): model runtime overview -- [src/graphon/graph_engine/layers/README.md](src/graphon/graph_engine/layers/README.md): +- [src/graphon/engine/layer/README.md](src/graphon/engine/layer/README.md): engine layer extension points -- [src/graphon/graph_engine/command_channels/README.md](src/graphon/graph_engine/command_channels/README.md): - local and distributed command channels +- [src/graphon/engine/command/README.md](src/graphon/engine/command/README.md): + command processing and local or distributed channels ## Development diff --git a/examples/slim_llm/code.py b/examples/slim_llm/code.py index 5caafd75..8c471f63 100644 --- a/examples/slim_llm/code.py +++ b/examples/slim_llm/code.py @@ -20,12 +20,11 @@ use_local_slim_binary, ) from graphon.dsl.slim import SlimLLM -from graphon.entities.graph_init_params import GraphInitParams +from graphon.engine import Engine +from graphon.entities.graph_init_params import InitParams from graphon.file.enums import FileType from graphon.file.models import File from graphon.graph.graph import Graph -from graphon.graph_engine.command_channels import InMemoryChannel -from graphon.graph_engine.graph_engine import GraphEngine from graphon.model_runtime.entities.llm_entities import LLMMode from graphon.model_runtime.entities.message_entities import ( PromptMessage, @@ -42,7 +41,7 @@ from graphon.nodes.llm.entities import ContextConfig from graphon.nodes.start import StartNode from graphon.nodes.start.entities import StartNodeData -from graphon.runtime.graph_runtime_state import GraphRuntimeState +from graphon.runtime.graph_runtime_state import RuntimeState from graphon.runtime.variable_pool import VariablePool @@ -79,9 +78,13 @@ def run(query: str) -> str: use_local_slim_binary() credentials = load_credentials() workflow_id = "slim-llm-code-example" - graph_state = GraphRuntimeState(variable_pool=VariablePool(), start_at=time.time()) + graph_state = RuntimeState( + workflow_id=workflow_id, + variable_pool=VariablePool(), + start_at=time.time(), + ) graph_state.variable_pool.add(("start", "query"), query) - graph_init = GraphInitParams( + graph_init = InitParams( workflow_id=workflow_id, graph_config={"nodes": [], "edges": []}, run_context={}, @@ -100,11 +103,9 @@ def run(query: str) -> str: parameters={}, ), ) - engine = GraphEngine( - workflow_id=workflow_id, + engine = Engine( graph=graph, graph_runtime_state=graph_state, - command_channel=InMemoryChannel(), ) list(engine.run()) @@ -117,8 +118,8 @@ def run(query: str) -> str: def build_graph( *, - graph_init: GraphInitParams, - graph_state: GraphRuntimeState, + graph_init: InitParams, + graph_state: RuntimeState, llm: SlimLLM, ) -> Graph: start = StartNode( diff --git a/examples/slim_llm/dsl.py b/examples/slim_llm/dsl.py index e7241121..3ebd4e56 100644 --- a/examples/slim_llm/dsl.py +++ b/examples/slim_llm/dsl.py @@ -15,13 +15,13 @@ use_local_slim_binary, ) from graphon.dsl import loads -from graphon.filters import ( - GraphEventFilterContext, +from graphon.engine.filter import ( + EngineEventFilterContext, ResponseStreamFilter, - filter_graph_events, + filter_engine_events, ) -from graphon.graph_events.graph import GraphRunSucceededEvent -from graphon.graph_events.node import NodeRunStreamChunkEvent +from graphon.engine_events.graph import GraphRunSucceededEvent +from graphon.engine_events.node import NodeRunStreamChunkEvent def run( @@ -37,9 +37,9 @@ def run( start_inputs={"query": query}, ) - events = filter_graph_events( + events = filter_engine_events( engine.run(), - context=GraphEventFilterContext.from_engine(engine), + context=EngineEventFilterContext.from_engine(engine), filters=[ResponseStreamFilter()], ) final_event: GraphRunSucceededEvent | None = None diff --git a/src/graphon/dsl/importer.py b/src/graphon/dsl/importer.py index 3ae6a1f8..f89dbdd0 100644 --- a/src/graphon/dsl/importer.py +++ b/src/graphon/dsl/importer.py @@ -10,15 +10,14 @@ import yaml from pydantic import ValidationError -from graphon.entities.graph_init_params import GraphInitParams +from graphon.engine import Engine +from graphon.engine.command import CommandChannel +from graphon.engine.container_handler import ContainerHandlerFactory +from graphon.entities.graph_init_params import InitParams from graphon.enums import BuiltinNodeTypes from graphon.graph.graph import Graph from graphon.graph.validation import GraphValidationError -from graphon.graph_engine.command_channels import CommandChannel, InMemoryChannel -from graphon.graph_engine.config import GraphEngineConfig -from graphon.graph_engine.container_handlers import ContainerHandlerFactory -from graphon.graph_engine.graph_engine import GraphEngine -from graphon.runtime.graph_runtime_state import GraphRuntimeState +from graphon.runtime.graph_runtime_state import RuntimeState from graphon.runtime.variable_pool import VariablePool from .entities import ( @@ -94,9 +93,9 @@ def loads( run_context: Mapping[str, Any] | None = None, start_inputs: Mapping[str, Any] | None = None, command_channel: CommandChannel | None = None, - config: GraphEngineConfig | None = None, + workers: int = 5, container_handler_factories: Sequence[ContainerHandlerFactory] = (), -) -> GraphEngine: +) -> Engine: plan = inspect(dsl, source_kind=source_kind) if plan.load_status == LoadStatus.UNSUPPORTED: raise _dsl_error( @@ -128,18 +127,18 @@ def loads( run_context=run_context or {}, start_inputs=start_inputs or {}, ) - graph_init_params = GraphInitParams( + graph_init_params = InitParams( workflow_id=workflow_id, graph_config=graph_config, run_context=run_context or {}, call_depth=0, ) - graph_runtime_state = GraphRuntimeState( + graph_runtime_state = RuntimeState( variable_pool=variable_pool, start_at=time.time(), + workflow_id=workflow_id, ) parsed_credentials = _parse_credentials(credentials) - engine_config = config or GraphEngineConfig() node_factory = SlimDslNodeFactory( graph_config=graph_config, graph_init_params=graph_init_params, @@ -170,12 +169,11 @@ def loads( kind=plan.document.kind, ) from error - return GraphEngine( - workflow_id=workflow_id, + return Engine( graph=graph, graph_runtime_state=graph_runtime_state, - command_channel=command_channel or InMemoryChannel(), - config=engine_config, + command_channel=command_channel, + workers=workers, container_handler_factories=container_handler_factories, ) diff --git a/src/graphon/dsl/node_factory.py b/src/graphon/dsl/node_factory.py index 63aa2c4a..3a5c6a69 100644 --- a/src/graphon/dsl/node_factory.py +++ b/src/graphon/dsl/node_factory.py @@ -66,7 +66,7 @@ from graphon.nodes.variable_assigner.v2.node import ( VariableAssignerNode as VariableAssignerNodeV2, ) -from graphon.runtime.graph_runtime_state import GraphRuntimeState +from graphon.runtime.graph_runtime_state import RuntimeState from graphon.template_rendering import Jinja2TemplateRenderer, TemplateRenderError from .code_runtime import SandboxCodeExecutor @@ -418,7 +418,7 @@ class _NodeBuildRequest: class SlimDslNodeFactory: graph_config: Mapping[str, Any] graph_init_params: Any - graph_runtime_state: GraphRuntimeState + graph_runtime_state: RuntimeState credentials: DslCredentials dependencies: list[DslDependency] slim_client_config: SlimClientConfig = field(init=False) @@ -442,7 +442,7 @@ def __post_init__(self) -> None: def with_runtime_state( self, - graph_runtime_state: GraphRuntimeState, + graph_runtime_state: RuntimeState, ) -> SlimDslNodeFactory: return replace(self, graph_runtime_state=graph_runtime_state) diff --git a/src/graphon/engine/__init__.py b/src/graphon/engine/__init__.py new file mode 100644 index 00000000..17243b35 --- /dev/null +++ b/src/graphon/engine/__init__.py @@ -0,0 +1,3 @@ +from .engine import Engine + +__all__ = ["Engine"] diff --git a/src/graphon/graph_engine/command_channels/README.md b/src/graphon/engine/command/README.md similarity index 52% rename from src/graphon/graph_engine/command_channels/README.md rename to src/graphon/engine/command/README.md index d0c693a2..063c76e5 100644 --- a/src/graphon/graph_engine/command_channels/README.md +++ b/src/graphon/engine/command/README.md @@ -1,9 +1,17 @@ -# Command Channels +# Commands -Channel implementations for external workflow control. +Command processing and channels for external workflow control. + +The supported command union contains `AbortCommand`, `PauseCommand`, and +`UpdateVariablesCommand`. Their `command_type` literals are used to deserialize +commands received from distributed channels. ## Components +### CommandProcessor + +Polls a command channel and applies commands to the current graph execution. + ### InMemoryChannel Thread-safe in-memory queue for single-process deployments. @@ -21,9 +29,11 @@ Redis-based queue for distributed deployments. ## Usage ```python +from graphon.engine.command import AbortCommand, InMemoryChannel, RedisChannel + # Local execution channel = InMemoryChannel() -channel.send_command(AbortCommand(graph_id="workflow-123")) +channel.send_command(AbortCommand(reason="stop")) # Distributed execution redis_channel = RedisChannel( diff --git a/src/graphon/engine/command/__init__.py b/src/graphon/engine/command/__init__.py new file mode 100644 index 00000000..bdce0a12 --- /dev/null +++ b/src/graphon/engine/command/__init__.py @@ -0,0 +1,23 @@ +"""Engine command communication and processing.""" + +from .builtin.in_memory import InMemoryChannel +from .builtin.redis import RedisChannel +from .entities import ( + AbortCommand, + Command, + PauseCommand, + UpdateVariablesCommand, +) +from .processor import CommandProcessor +from .protocol import CommandChannel + +__all__ = [ + "AbortCommand", + "Command", + "CommandChannel", + "CommandProcessor", + "InMemoryChannel", + "PauseCommand", + "RedisChannel", + "UpdateVariablesCommand", +] diff --git a/src/graphon/engine/command/builtin/__init__.py b/src/graphon/engine/command/builtin/__init__.py new file mode 100644 index 00000000..9cc4c897 --- /dev/null +++ b/src/graphon/engine/command/builtin/__init__.py @@ -0,0 +1 @@ +"""Implementation modules for the command channels exported by the parent package.""" diff --git a/src/graphon/graph_engine/command_channels/in_memory_channel.py b/src/graphon/engine/command/builtin/in_memory.py similarity index 76% rename from src/graphon/graph_engine/command_channels/in_memory_channel.py rename to src/graphon/engine/command/builtin/in_memory.py index 557100e6..bf7a0417 100644 --- a/src/graphon/graph_engine/command_channels/in_memory_channel.py +++ b/src/graphon/engine/command/builtin/in_memory.py @@ -7,29 +7,29 @@ from queue import Empty, Queue from typing import final -from ..entities.commands import GraphEngineCommand +from ..entities import Command @final class InMemoryChannel: """In-memory command channel implementation using a thread-safe queue. - Each instance is dedicated to a single GraphEngine/workflow execution. + Each instance is dedicated to a single Engine/workflow execution. Suitable for local development, testing, and single-instance deployments. """ def __init__(self) -> None: """Initialize the in-memory channel with a single queue.""" - self._queue: Queue[GraphEngineCommand] = Queue() + self._queue: Queue[Command] = Queue() - def fetch_commands(self) -> list[GraphEngineCommand]: + def fetch_commands(self) -> list[Command]: """Fetch all pending commands from the queue. Returns: List of pending commands (drains the queue) """ - commands: list[GraphEngineCommand] = [] + commands: list[Command] = [] # Drain all available commands from the queue while not self._queue.empty(): @@ -41,7 +41,7 @@ def fetch_commands(self) -> list[GraphEngineCommand]: return commands - def send_command(self, command: GraphEngineCommand) -> None: + def send_command(self, command: Command) -> None: """Send a command to this channel's queue. Args: diff --git a/src/graphon/graph_engine/command_channels/redis_channel.py b/src/graphon/engine/command/builtin/redis.py similarity index 58% rename from src/graphon/graph_engine/command_channels/redis_channel.py rename to src/graphon/engine/command/builtin/redis.py index f55b0457..95f4a47d 100644 --- a/src/graphon/graph_engine/command_channels/redis_channel.py +++ b/src/graphon/engine/command/builtin/redis.py @@ -6,54 +6,72 @@ """ import json -from abc import abstractmethod from contextlib import AbstractContextManager -from typing import Any, Protocol, final +from typing import Any, Protocol, cast, final -from ..entities.commands import ( - AbortCommand, - CommandType, - GraphEngineCommand, - PauseCommand, - UpdateVariablesCommand, -) +from pydantic import TypeAdapter -_COMMAND_MODEL_BY_TYPE: dict[CommandType, type[GraphEngineCommand]] = { - CommandType.ABORT: AbortCommand, - CommandType.PAUSE: PauseCommand, - CommandType.UPDATE_VARIABLES: UpdateVariablesCommand, -} +from ..entities import Command + +_COMMAND_ADAPTER = TypeAdapter(Command) + + +def _migrate_command_payload(data: object) -> object: + """Normalize the legacy Redis wire shape for variable update commands. + + Older producers wrapped each serialized variable as ``{"value": variable}``. + The public wrapper type no longer exists, but Redis commands can remain queued + for up to one hour during a rolling deployment. This boundary-only migration + unwraps those list items without reintroducing a Python compatibility type or + changing already-current payloads. + + Args: + data: Decoded command JSON received from Redis. Non-object JSON is + returned unchanged so command validation can reject it normally. + + Returns: + The original mapping when no migration is needed, otherwise a shallow copy + whose ``updates`` list contains the serialized variables directly. + + """ + if not isinstance(data, dict): + return data + + command_data = cast(dict[str, Any], data) + updates = command_data.get("updates") + if command_data.get("command_type") != "update_variables" or not isinstance( + updates, + list, + ): + return command_data + migrated = [ + update["value"] + if isinstance(update, dict) and set(update) == {"value"} + else update + for update in updates + ] + return ( + command_data if migrated == updates else {**command_data, "updates": migrated} + ) class RedisPipelineProtocol(Protocol): """Minimal Redis pipeline contract used by the command channel.""" - @abstractmethod def lrange(self, name: str, start: int, end: int) -> Any: ... - @abstractmethod def delete(self, *names: str) -> Any: ... - @abstractmethod def execute(self) -> list[Any]: ... - @abstractmethod def rpush(self, name: str, *values: str) -> Any: ... - @abstractmethod def expire(self, name: str, time: int) -> Any: ... - @abstractmethod - def set(self, name: str, value: str, ex: int | None = None) -> Any: ... - - @abstractmethod - def get(self, name: str) -> Any: ... - class RedisClientProtocol(Protocol): """Redis client contract required by the command channel.""" - @abstractmethod def pipeline(self) -> AbstractContextManager[RedisPipelineProtocol]: ... @@ -82,19 +100,15 @@ def __init__( self._redis = redis_client self._key = channel_key self._command_ttl = command_ttl - self._pending_key = f"{channel_key}:pending" - def fetch_commands(self) -> list[GraphEngineCommand]: + def fetch_commands(self) -> list[Command]: """Fetch all pending commands from Redis. Returns: List of pending commands (drains the Redis list) """ - if not self._has_pending_commands(): - return [] - - commands: list[GraphEngineCommand] = [] + commands: list[Command] = [] # Use pipeline for atomic operations with self._redis.pipeline() as pipe: @@ -108,7 +122,7 @@ def fetch_commands(self) -> list[GraphEngineCommand]: for command_json in results[0]: try: command_data = json.loads(command_json) - command = self._deserialize_command(command_data) + command = self.deserialize_command(command_data) if command: commands.append(command) except (json.JSONDecodeError, ValueError): @@ -117,7 +131,7 @@ def fetch_commands(self) -> list[GraphEngineCommand]: return commands - def send_command(self, command: GraphEngineCommand) -> None: + def send_command(self, command: Command) -> None: """Send a command to Redis. Args: @@ -130,48 +144,22 @@ def send_command(self, command: GraphEngineCommand) -> None: with self._redis.pipeline() as pipe: pipe.rpush(self._key, command_json) pipe.expire(self._key, self._command_ttl) - pipe.set(self._pending_key, "1", ex=self._command_ttl) pipe.execute() def deserialize_command( self, - data: dict[str, Any], - ) -> GraphEngineCommand | None: - """Deserialize a command payload into a typed command model.""" - return self._deserialize_command(data) - - def _deserialize_command(self, data: dict[str, Any]) -> GraphEngineCommand | None: + data: object, + ) -> Command | None: """Deserialize a command from dictionary data. Args: - data: Command data dictionary + data: Decoded JSON value received from the command queue. Returns: Deserialized command or None if invalid """ - command_type_value = data.get("command_type") - if not isinstance(command_type_value, str): - return None - try: - command_type = CommandType(command_type_value) - command_model = _COMMAND_MODEL_BY_TYPE.get(command_type, GraphEngineCommand) - return command_model.model_validate(data) - + return _COMMAND_ADAPTER.validate_python(_migrate_command_payload(data)) except (ValueError, TypeError): return None - - def _has_pending_commands(self) -> bool: - """Check and consume the pending marker to avoid unnecessary list reads. - - Returns: - True if commands should be fetched from Redis. - - """ - with self._redis.pipeline() as pipe: - pipe.get(self._pending_key) - pipe.delete(self._pending_key) - pending_value, _ = pipe.execute() - - return pending_value is not None diff --git a/src/graphon/engine/command/entities.py b/src/graphon/engine/command/entities.py new file mode 100644 index 00000000..6f3178cf --- /dev/null +++ b/src/graphon/engine/command/entities.py @@ -0,0 +1,42 @@ +"""Engine command entities for external control. + +This module defines command types that can be sent to a running Engine +instance to control its execution flow. +""" + +from collections.abc import Sequence +from typing import Annotated, Literal + +from pydantic import BaseModel, Field + +from graphon.variables.variables import Variable + + +class AbortCommand(BaseModel): + """Command to abort a running workflow execution.""" + + command_type: Literal["abort"] = "abort" + reason: str | None = Field(default=None, description="Optional reason for abort") + + +class PauseCommand(BaseModel): + """Command to pause a running workflow execution.""" + + command_type: Literal["pause"] = "pause" + reason: str = Field(default="unknown reason", description="reason for pause") + + +class UpdateVariablesCommand(BaseModel): + """Command to update a group of variables in the variable pool.""" + + command_type: Literal["update_variables"] = "update_variables" + updates: Sequence[Variable] = Field( + default_factory=list, + description="Variable updates", + ) + + +type Command = Annotated[ + AbortCommand | PauseCommand | UpdateVariablesCommand, + Field(discriminator="command_type"), +] diff --git a/src/graphon/engine/command/processor.py b/src/graphon/engine/command/processor.py new file mode 100644 index 00000000..12ea9a3c --- /dev/null +++ b/src/graphon/engine/command/processor.py @@ -0,0 +1,109 @@ +"""Main command processor for handling external commands.""" + +import logging +from typing import final + +from graphon.entities.pause_reason import SchedulingPause +from graphon.runtime.execution import GraphExecution +from graphon.runtime.variable_pool import VariablePool + +from .entities import ( + AbortCommand, + Command, + PauseCommand, + UpdateVariablesCommand, +) +from .protocol import CommandChannel + +logger = logging.getLogger(__name__) + + +@final +class CommandProcessor: + """Processes external commands sent to the engine. + + This polls the command channel and applies each supported command directly. + """ + + def __init__( + self, + command_channel: CommandChannel, + graph_execution: GraphExecution, + variable_pool: VariablePool, + ) -> None: + """Initialize the command processor. + + Args: + command_channel: Channel for receiving commands + graph_execution: Graph execution aggregate + variable_pool: Runtime variables updated by external commands + + """ + self._command_channel = command_channel + self._graph_execution = graph_execution + self._variable_pool = variable_pool + + def process_commands(self) -> None: + """Check for and process any pending commands.""" + try: + commands = self._command_channel.fetch_commands() + except Exception: + logger.exception("Error processing commands") + return + + for command in commands: + try: + self._handle_command(command) + except Exception: + logger.exception( + "Error handling command %s", + command.__class__.__name__, + ) + + def _handle_command(self, command: Command) -> None: + """Handle a single command. + + Args: + command: The command to handle + + """ + match command: + case AbortCommand(): + logger.debug( + "Aborting workflow %s: %s", + self._graph_execution.workflow_id, + command.reason, + ) + self._graph_execution.abort( + command.reason or "User requested abort", + ) + case PauseCommand(): + logger.debug( + "Pausing workflow %s: %s", + self._graph_execution.workflow_id, + command.reason, + ) + self._graph_execution.pause( + SchedulingPause(message=command.reason), + ) + case UpdateVariablesCommand(): + for variable in command.updates: + try: + self._variable_pool.add(variable.selector, variable) + logger.debug( + "Updated variable %s for workflow %s", + variable.selector, + self._graph_execution.workflow_id, + ) + except ValueError as exc: + logger.warning( + "Skipping invalid variable selector %s for workflow %s: %s", + getattr(variable, "selector", None), + self._graph_execution.workflow_id, + exc, + ) + case _: + logger.warning( + "Unsupported command: %s", + command.__class__.__name__, + ) diff --git a/src/graphon/engine/command/protocol.py b/src/graphon/engine/command/protocol.py new file mode 100644 index 00000000..2eee6a3d --- /dev/null +++ b/src/graphon/engine/command/protocol.py @@ -0,0 +1,39 @@ +"""CommandChannel protocol for Engine command communication. + +This protocol defines the interface for sending and receiving commands +to/from an Engine instance, supporting both local and distributed scenarios. +""" + +from typing import Protocol + +from .entities import Command + + +class CommandChannel(Protocol): + """Protocol for bidirectional command communication with Engine. + + Since each Engine instance processes only one workflow execution, + this channel is dedicated to that single execution. + """ + + def fetch_commands(self) -> list[Command]: + """Fetch pending commands for this Engine instance. + + Called by Engine to poll for commands that need to be processed. + + Returns: + List of pending commands (may be empty) + + """ + ... + + def send_command(self, command: Command) -> None: + """Send a command to be processed by this Engine instance. + + Called by external systems to send control commands to the running workflow. + + Args: + command: The command to send + + """ + ... diff --git a/src/graphon/engine/container_handler/__init__.py b/src/graphon/engine/container_handler/__init__.py new file mode 100644 index 00000000..5ff8e3b9 --- /dev/null +++ b/src/graphon/engine/container_handler/__init__.py @@ -0,0 +1,13 @@ +"""Container handler contracts and built-in implementations.""" + +from .builtin.iteration import IterationContainerHandler +from .builtin.loop import LoopContainerHandler +from .container_handler_factory import ContainerHandlerFactory +from .protocol import ContainerHandler + +__all__ = [ + "ContainerHandler", + "ContainerHandlerFactory", + "IterationContainerHandler", + "LoopContainerHandler", +] diff --git a/src/graphon/engine/container_handler/builtin/__init__.py b/src/graphon/engine/container_handler/builtin/__init__.py new file mode 100644 index 00000000..e383c6cc --- /dev/null +++ b/src/graphon/engine/container_handler/builtin/__init__.py @@ -0,0 +1 @@ +"""Implementation modules for container handlers exported by the parent package.""" diff --git a/src/graphon/graph_engine/iteration_container_handler.py b/src/graphon/engine/container_handler/builtin/iteration.py similarity index 85% rename from src/graphon/graph_engine/iteration_container_handler.py rename to src/graphon/engine/container_handler/builtin/iteration.py index c2294581..33200ddb 100644 --- a/src/graphon/graph_engine/iteration_container_handler.py +++ b/src/graphon/engine/container_handler/builtin/iteration.py @@ -1,16 +1,18 @@ +"""Built-in iteration container handler.""" + from __future__ import annotations from datetime import UTC, datetime from typing import cast, final +from graphon.engine_events.base import NodeEvent +from graphon.engine_events.node import NodeRunFailedEvent from graphon.enums import ( BuiltinNodeTypes, ErrorHandleMode, WorkflowNodeExecutionMetadataKey, WorkflowNodeExecutionStatus, ) -from graphon.graph_events.base import GraphNodeEventBase -from graphon.graph_events.node import NodeRunFailedEvent from graphon.nodes.container_effects import ( ContainerAwaitRequest, ContainerExecutionResult, @@ -24,11 +26,12 @@ IterationFrameState, IterationRunState, ) -from graphon.runtime.graph_runtime_state import GraphRuntimeState +from graphon.runtime.execution import ROOT_FRAME_ID +from graphon.runtime.graph_runtime_state import RuntimeState from graphon.variables.segments import NoneSegment, SerializableSegment -from .frames import ExecutionFrame, FrameRegistry -from .ready_queue import ROOT_FRAME_ID, ResumeTask +from ...frame import ExecutionFrame, FrameRegistry +from ...ready_queue import ResumeTask @final @@ -51,17 +54,34 @@ def restore_frame(self, frame_state: ContainerFrameState) -> None: f"iteration frame {frame_state.frame_id} requires a local variable pool" ) raise TypeError(msg) - self._frame_registry.materialize_child_frame_from_state( - frame_state, + run_state = self._iteration_run(frame_state.parent_invocation_id) + self._frame_registry.restore_child( + frame_id=frame_state.frame_id, + parent_frame_id=run_state.frame_id, + container_id=run_state.node_id, + root_node_id=frame_state.root_node_id, + runtime_data=frame_state.runtime_data, variable_pool=variable_pool.model_copy(deep=True), ) - def start_await( + def handle_request( self, *, invocation_id: str, request: ContainerAwaitRequest, ) -> None: + """Schedule iteration frames requested by a suspended Iteration node. + + The request is validated against the handler type, scheduled indexes are + recorded, and as many child frames as the configured parallel capacity + allows are started. A previously terminal failure resumes the parent + immediately instead. + + Raises: + TypeError: If the request or its stored run state is not an + iteration value. + + """ if not isinstance(request, IterationFrameRequest): msg = f"iteration handler cannot handle {type(request).__name__}" raise TypeError(msg) @@ -70,7 +90,7 @@ def start_await( run_state = self._put_run_state( run_state.model_copy(update={"resume_pending": False}), ) - parent_frame = self._frame_registry.get(run_state.frame_id) + parent_frame = self._frame_registry[run_state.frame_id] if self._finish_failed_iteration_if_ready( parent_frame=parent_frame, run_state=run_state, @@ -103,12 +123,10 @@ def prepare_frame_event( self, *, frame: ExecutionFrame, - event: GraphNodeEventBase, + event: NodeEvent, ) -> None: frame_state = self._iteration_frame(frame.frame_id) run_state = self._iteration_run(frame_state.parent_invocation_id) - if event.in_iteration_id is None: - event.in_iteration_id = run_state.node_id iteration_metadata = { WorkflowNodeExecutionMetadataKey.ITERATION_ID: run_state.node_id, WorkflowNodeExecutionMetadataKey.ITERATION_INDEX: frame_state.index, @@ -120,11 +138,12 @@ def prepare_frame_event( **iteration_metadata, } - def should_collect( + def should_emit( self, *, - event: GraphNodeEventBase, + event: NodeEvent, ) -> bool: + """Hide the synthetic Iteration Start event from outer observers.""" return event.node_type != BuiltinNodeTypes.ITERATION_START def record_frame_failure( @@ -140,8 +159,14 @@ def record_frame_failure( ), ) - def complete_frame(self, frame: ExecutionFrame) -> None: - if not frame.state_manager.is_execution_complete(): + def complete_frame_if_ready(self, frame: ExecutionFrame) -> None: + """Finalize and remove an iteration frame once all its nodes finish. + + Successful frames store their selected output; failed frames apply the + configured error policy. Frames that still have unfinished nodes are + left untouched because this hook runs after every node event. + """ + if not frame.scheduler.is_execution_complete(): return root_runtime_state = self._root_runtime_state() @@ -160,7 +185,7 @@ def _complete_ready_iteration_frame( frame_state: IterationFrameState, ) -> None: run_state = self._iteration_run(frame_state.parent_invocation_id) - parent_frame = self._frame_registry.get(run_state.frame_id) + parent_frame = self._frame_registry[run_state.frame_id] if frame_state.errors: self._complete_failed_iteration_frame( frame=frame, @@ -171,7 +196,7 @@ def _complete_ready_iteration_frame( ) return - result = frame.graph_runtime_state.variable_pool.get( + result = frame.state.variable_pool.get( run_state.output_selector, ) output = NoneSegment() if result is None else cast(SerializableSegment, result) @@ -245,8 +270,8 @@ def _complete_iteration_step( if store_output: outputs[str(frame_state.index)] = output if not frame_state.errors and frame_state.index == len(run_state.items) - 1: - parent_frame.graph_runtime_state.merge_response_outputs( - frame.graph_runtime_state.outputs, + parent_frame.state.merge_response_outputs( + frame.state.outputs, ) return self._put_run_state( @@ -255,7 +280,7 @@ def _complete_iteration_step( "outputs": outputs, "duration_map": duration_map, "usage": run_state.usage.plus( - frame.graph_runtime_state.llm_usage, + frame.state.llm_usage, ), "completed_count": run_state.completed_count + 1, }, @@ -275,7 +300,7 @@ def _continue_or_complete_iteration( return if run_state.completed_count >= len(run_state.items): self._enqueue_container_result( - runtime_state=parent_frame.graph_runtime_state, + runtime_state=parent_frame.state, invocation_id=run_state.invocation_id, result=self._build_iteration_result( run_state=run_state, @@ -307,7 +332,7 @@ def _finish_failed_iteration_if_ready( run_state.model_copy(update={"resume_pending": True}), ) self._enqueue_container_result( - runtime_state=parent_frame.graph_runtime_state, + runtime_state=parent_frame.state, invocation_id=run_state.invocation_id, result=self._build_iteration_result( run_state=run_state, @@ -340,7 +365,7 @@ def _request_iteration_frames( run_state.model_copy(update={"resume_pending": True}), ) self._enqueue_container_result( - runtime_state=parent_frame.graph_runtime_state, + runtime_state=parent_frame.state, invocation_id=run_state.invocation_id, result=IterationFrameRequest( items=run_state.items, @@ -360,25 +385,18 @@ def _start_iteration_frame( run_state: IterationRunState, index: int, ) -> None: - variable_pool = parent_frame.graph_runtime_state.variable_pool.model_copy( + variable_pool = parent_frame.state.variable_pool.model_copy( deep=True, ) variable_pool.add([run_state.node_id, "index"], index) variable_pool.add([run_state.node_id, "item"], run_state.items[index]) - child_runtime_state = GraphRuntimeState( - variable_pool=variable_pool, - start_at=parent_frame.graph_runtime_state.start_at, - ready_queue=parent_frame.graph_runtime_state.ready_queue, - deferred_ready_queue=( - parent_frame.graph_runtime_state.deferred_ready_queue - ), - graph_execution=parent_frame.graph_runtime_state.graph_execution, - ) child_frame_id = f"{run_state.invocation_id}:iteration:{index}" - child_frame = self._frame_registry.materialize_child_frame( + child_frame = self._frame_registry.create_child( frame_id=child_frame_id, + parent_frame_id=parent_frame.frame_id, + container_id=run_state.node_id, root_node_id=run_state.root_node_id, - graph_runtime_state=child_runtime_state, + variable_pool=variable_pool, ) self._root_runtime_state().put_container_frame( IterationFrameState( @@ -387,12 +405,12 @@ def _start_iteration_frame( root_node_id=run_state.root_node_id, index=index, started_at=datetime.now(UTC).replace(tzinfo=None), - runtime_data=child_frame.graph_runtime_state.snapshot_frame( + runtime_data=child_frame.state.snapshot_frame( copy_variable_pool=False, ), ), ) - child_frame.state_manager.enqueue_node(run_state.root_node_id) + child_frame.scheduler.enqueue_node(run_state.root_node_id) def _build_iteration_result( self, @@ -472,7 +490,7 @@ def _ordered_iteration_outputs( def _enqueue_container_result( self, *, - runtime_state: GraphRuntimeState, + runtime_state: RuntimeState, invocation_id: str, result: ContainerExecutionResult | IterationFrameRequest, ) -> None: @@ -511,5 +529,5 @@ def _put_run_state(self, run_state: IterationRunState) -> IterationRunState: self._root_runtime_state().put_container_run(run_state) return run_state - def _root_runtime_state(self) -> GraphRuntimeState: - return self._frame_registry.get(ROOT_FRAME_ID).graph_runtime_state + def _root_runtime_state(self) -> RuntimeState: + return self._frame_registry[ROOT_FRAME_ID].state diff --git a/src/graphon/graph_engine/loop_container_handler.py b/src/graphon/engine/container_handler/builtin/loop.py similarity index 78% rename from src/graphon/graph_engine/loop_container_handler.py rename to src/graphon/engine/container_handler/builtin/loop.py index 59b17235..0746bd19 100644 --- a/src/graphon/graph_engine/loop_container_handler.py +++ b/src/graphon/engine/container_handler/builtin/loop.py @@ -1,16 +1,18 @@ +"""Built-in loop container handler.""" + from __future__ import annotations import contextlib from datetime import UTC, datetime from typing import final +from graphon.engine_events.base import NodeEvent +from graphon.engine_events.node import NodeRunFailedEvent, NodeRunSucceededEvent from graphon.enums import ( BuiltinNodeTypes, WorkflowNodeExecutionMetadataKey, WorkflowNodeExecutionStatus, ) -from graphon.graph_events.base import GraphNodeEventBase -from graphon.graph_events.node import NodeRunFailedEvent, NodeRunSucceededEvent from graphon.nodes.container_effects import ( ContainerAwaitRequest, ContainerExecutionResult, @@ -24,12 +26,13 @@ LoopFrameState, LoopRunState, ) -from graphon.runtime.graph_runtime_state import GraphRuntimeState +from graphon.runtime.execution import ROOT_FRAME_ID +from graphon.runtime.graph_runtime_state import RuntimeState from graphon.utils.condition.processor import ConditionProcessor from graphon.variables.segments import SerializableSegment -from .frames import ExecutionFrame, FrameRegistry -from .ready_queue import ROOT_FRAME_ID, ResumeTask +from ...frame import ExecutionFrame, FrameRegistry +from ...ready_queue import ResumeTask @final @@ -53,18 +56,44 @@ def restore_frame(self, frame_state: ContainerFrameState) -> None: if not isinstance(run_state, LoopRunState): msg = f"loop frame cannot belong to {run_state.kind} run" raise TypeError(msg) - parent_frame = self._frame_registry.get(run_state.frame_id) - self._frame_registry.materialize_child_frame_from_state( - frame_state, - variable_pool=parent_frame.graph_runtime_state.variable_pool, + variable_pool = frame_state.runtime_data.variable_pool + inherited_parent_pool = isinstance(variable_pool, str) + if inherited_parent_pool: + variable_pool = self._frame_registry[run_state.frame_id].state.variable_pool + frame = self._frame_registry.restore_child( + frame_id=frame_state.frame_id, + parent_frame_id=run_state.frame_id, + container_id=run_state.node_id, + root_node_id=frame_state.root_node_id, + runtime_data=frame_state.runtime_data, + variable_pool=variable_pool.model_copy(deep=True), ) + if inherited_parent_pool: + root_runtime_state.put_container_frame( + frame_state.model_copy( + update={ + "runtime_data": frame.state.snapshot_frame(), + }, + ), + ) - def start_await( + def handle_request( self, *, invocation_id: str, request: ContainerAwaitRequest, ) -> None: + """Start the next loop frame requested by a suspended Loop node. + + The request and persisted run state are validated before evaluating the + break condition against the parent frame. A satisfied break resumes the + parent with a terminal result; otherwise a new scoped child is started. + + Raises: + TypeError: If the request, stored run, or owning graph node is not + a loop value. + + """ if not isinstance(request, LoopFrameRequest): msg = f"loop handler cannot handle {type(request).__name__}" raise TypeError(msg) @@ -73,7 +102,7 @@ def start_await( if not isinstance(run_state, LoopRunState): msg = f"loop handler cannot continue {run_state.kind} run" raise TypeError(msg) - parent_frame = self._frame_registry.get(run_state.frame_id) + parent_frame = self._frame_registry[run_state.frame_id] node = parent_frame.graph.nodes[run_state.node_id] if not isinstance(node, LoopNode): msg = f"node {run_state.node_id} cannot handle loop await requests" @@ -88,7 +117,7 @@ def start_await( ) self._root_runtime_state().put_container_run(run_state) self._enqueue_container_result( - runtime_state=parent_frame.graph_runtime_state, + runtime_state=parent_frame.state, invocation_id=run_state.invocation_id, result=self._complete_loop( run_state=run_state, @@ -107,7 +136,7 @@ def prepare_frame_event( self, *, frame: ExecutionFrame, - event: GraphNodeEventBase, + event: NodeEvent, ) -> None: root_runtime_state = self._root_runtime_state() frame_state = root_runtime_state.get_container_frame(frame.frame_id) @@ -120,8 +149,6 @@ def prepare_frame_event( if not isinstance(run_state, LoopRunState): msg = f"loop frame cannot belong to {run_state.kind} run" raise TypeError(msg) - if event.in_loop_id is None: - event.in_loop_id = run_state.node_id loop_metadata = { WorkflowNodeExecutionMetadataKey.LOOP_ID: run_state.node_id, WorkflowNodeExecutionMetadataKey.LOOP_INDEX: frame_state.index, @@ -143,11 +170,12 @@ def prepare_frame_event( ), ) - def should_collect( + def should_emit( self, *, - event: GraphNodeEventBase, + event: NodeEvent, ) -> bool: + """Hide the synthetic Loop Start event from outer observers.""" return event.node_type != BuiltinNodeTypes.LOOP_START def record_frame_failure( @@ -166,8 +194,19 @@ def record_frame_failure( ), ) - def complete_frame(self, frame: ExecutionFrame) -> None: - if not frame.state_manager.is_execution_complete(): + def complete_frame_if_ready(self, frame: ExecutionFrame) -> None: + """Finalize and remove a loop frame once all its nodes finish. + + The completed step updates loop variables, usage, and duration before + either scheduling the next round or resuming the parent. Frames that + still have unfinished nodes are left untouched. + + Raises: + TypeError: If restored container state has inconsistent types. + ValueError: If a completed child lacks a required loop variable. + + """ + if not frame.scheduler.is_execution_complete(): return root_runtime_state = self._root_runtime_state() @@ -186,9 +225,9 @@ def complete_frame(self, frame: ExecutionFrame) -> None: ) if not isinstance(run_state, LoopRunState): raise - parent_frame = self._frame_registry.get(run_state.frame_id) + parent_frame = self._frame_registry[run_state.frame_id] self._enqueue_container_result( - runtime_state=parent_frame.graph_runtime_state, + runtime_state=parent_frame.state, invocation_id=run_state.invocation_id, result=self._fail_loop(run_state=run_state, error=str(error)), ) @@ -209,7 +248,7 @@ def _complete_ready_loop_frame( if not isinstance(run_state, LoopRunState): msg = f"loop frame cannot complete {run_state.kind} run" raise TypeError(msg) - parent_frame = self._frame_registry.get(run_state.frame_id) + parent_frame = self._frame_registry[run_state.frame_id] node = parent_frame.graph.nodes[run_state.node_id] if not isinstance(node, LoopNode): msg = f"node {run_state.node_id} is not a loop" @@ -223,7 +262,7 @@ def _complete_ready_loop_frame( ) if frame_state.errors: self._enqueue_container_result( - runtime_state=parent_frame.graph_runtime_state, + runtime_state=parent_frame.state, invocation_id=run_state.invocation_id, result=self._fail_loop( run_state=run_state, @@ -234,7 +273,7 @@ def _complete_ready_loop_frame( if frame_state.reached_break or ( self._loop_break_conditions_reached( - frame=parent_frame, + frame=frame, node=node, suppress_errors=False, ) @@ -246,7 +285,7 @@ def _complete_ready_loop_frame( if run_state.reached_break or run_state.completed_count >= run_state.loop_count: self._enqueue_container_result( - runtime_state=parent_frame.graph_runtime_state, + runtime_state=parent_frame.state, invocation_id=run_state.invocation_id, result=self._complete_loop( run_state=run_state, @@ -256,7 +295,7 @@ def _complete_ready_loop_frame( return self._enqueue_container_result( - runtime_state=parent_frame.graph_runtime_state, + runtime_state=parent_frame.state, invocation_id=run_state.invocation_id, result=LoopFrameRequest( inputs=run_state.inputs, @@ -276,22 +315,22 @@ def _start_loop_frame( run_state: LoopRunState, request: LoopFrameRequest, ) -> None: - for node_id in request.loop_node_ids: - parent_frame.graph_runtime_state.variable_pool.remove([node_id]) - child_runtime_state = GraphRuntimeState( - variable_pool=parent_frame.graph_runtime_state.variable_pool, - start_at=parent_frame.graph_runtime_state.start_at, - ready_queue=parent_frame.graph_runtime_state.ready_queue, - deferred_ready_queue=( - parent_frame.graph_runtime_state.deferred_ready_queue - ), - graph_execution=parent_frame.graph_runtime_state.graph_execution, + variable_pool = parent_frame.state.variable_pool.model_copy( + deep=True, ) + for node_id in request.loop_node_ids: + variable_pool.remove([node_id]) + for key, selector in run_state.loop_variable_selectors.items(): + value = run_state.outputs.get(key) + if value is not None: + variable_pool.add(selector, value) child_frame_id = f"{run_state.invocation_id}:loop:{request.index}" - child_frame = self._frame_registry.materialize_child_frame( + child_frame = self._frame_registry.create_child( frame_id=child_frame_id, + parent_frame_id=parent_frame.frame_id, + container_id=run_state.node_id, root_node_id=request.root_node_id, - graph_runtime_state=child_runtime_state, + variable_pool=variable_pool, ) self._root_runtime_state().put_container_frame( LoopFrameState( @@ -300,12 +339,12 @@ def _start_loop_frame( root_node_id=request.root_node_id, index=request.index, started_at=datetime.now(UTC).replace(tzinfo=None), - runtime_data=child_frame.graph_runtime_state.snapshot_frame( - variable_pool_scope="parent", + runtime_data=child_frame.state.snapshot_frame( + copy_variable_pool=False, ), ), ) - child_frame.state_manager.enqueue_node(request.root_node_id) + child_frame.scheduler.enqueue_node(request.root_node_id) def _complete_loop_step( self, @@ -316,7 +355,7 @@ def _complete_loop_step( run_state: LoopRunState, ) -> LoopRunState: completed_count = run_state.completed_count + 1 - usage = run_state.usage.plus(frame.graph_runtime_state.llm_usage) + usage = run_state.usage.plus(frame.state.llm_usage) duration_map = dict(run_state.duration_map) loop_index = frame_state.index duration_map[str(loop_index)] = ( @@ -328,15 +367,15 @@ def _complete_loop_step( } loop_variable_values: dict[str, SerializableSegment] = {} for key, selector in run_state.loop_variable_selectors.items(): - segment = parent_frame.graph_runtime_state.variable_pool.get(selector) + segment = frame.state.variable_pool.get(selector) if segment is None: msg = f"loop variable {key} is missing" raise ValueError(msg) loop_variable_values[key] = build_container_value(segment) variable_map[str(loop_index)] = loop_variable_values if not frame_state.errors: - parent_frame.graph_runtime_state.merge_response_outputs( - frame.graph_runtime_state.outputs, + parent_frame.state.merge_response_outputs( + frame.state.outputs, ) outputs.update(loop_variable_values) outputs["loop_round"] = build_container_value(loop_index + 1) @@ -399,7 +438,7 @@ def _fail_loop( def _enqueue_container_result( self, *, - runtime_state: GraphRuntimeState, + runtime_state: RuntimeState, invocation_id: str, result: ContainerExecutionResult | LoopFrameRequest, ) -> None: @@ -441,7 +480,7 @@ def _loop_break_conditions_reached( if suppress_errors: with contextlib.suppress(ValueError): _, _, result = condition_processor.process_conditions( - variable_pool=frame.graph_runtime_state.variable_pool, + variable_pool=frame.state.variable_pool, conditions=node.node_data.break_conditions, operator=node.node_data.logical_operator, ) @@ -449,11 +488,11 @@ def _loop_break_conditions_reached( return False _, _, result = condition_processor.process_conditions( - variable_pool=frame.graph_runtime_state.variable_pool, + variable_pool=frame.state.variable_pool, conditions=node.node_data.break_conditions, operator=node.node_data.logical_operator, ) return result - def _root_runtime_state(self) -> GraphRuntimeState: - return self._frame_registry.get(ROOT_FRAME_ID).graph_runtime_state + def _root_runtime_state(self) -> RuntimeState: + return self._frame_registry[ROOT_FRAME_ID].state diff --git a/src/graphon/engine/container_handler/container_handler_factory.py b/src/graphon/engine/container_handler/container_handler_factory.py new file mode 100644 index 00000000..79b60dd3 --- /dev/null +++ b/src/graphon/engine/container_handler/container_handler_factory.py @@ -0,0 +1,8 @@ +"""Factory type for constructing container handlers.""" + +from collections.abc import Callable + +from ..frame import FrameRegistry +from .protocol import ContainerHandler + +type ContainerHandlerFactory = Callable[[FrameRegistry], ContainerHandler] diff --git a/src/graphon/engine/container_handler/protocol.py b/src/graphon/engine/container_handler/protocol.py new file mode 100644 index 00000000..e02aa400 --- /dev/null +++ b/src/graphon/engine/container_handler/protocol.py @@ -0,0 +1,62 @@ +"""Protocol implemented by every container handler.""" + +from __future__ import annotations + +from typing import Protocol + +from graphon.engine_events.base import NodeEvent +from graphon.engine_events.node import NodeRunFailedEvent +from graphon.enums import NodeType +from graphon.nodes.container_effects import ContainerAwaitRequest +from graphon.runtime.container_state import ContainerFrameState + +from ..frame import ExecutionFrame + + +class ContainerHandler(Protocol): + node_type: NodeType + + def restore_frame(self, frame_state: ContainerFrameState) -> None: ... + + def handle_request( + self, + *, + invocation_id: str, + request: ContainerAwaitRequest, + ) -> None: + """Handle a request emitted by a suspended container node. + + Implementations schedule or restore the child-frame work needed before + the container invocation can resume. + """ + ... + + def prepare_frame_event( + self, + *, + frame: ExecutionFrame, + event: NodeEvent, + ) -> None: ... + + def should_emit( + self, + *, + event: NodeEvent, + ) -> bool: + """Return whether a child-frame event should leave the container.""" + ... + + def record_frame_failure( + self, + *, + frame: ExecutionFrame, + event: NodeRunFailedEvent, + ) -> None: ... + + def complete_frame_if_ready(self, frame: ExecutionFrame) -> None: + """Finalize a child frame when its scheduler reports completion. + + The hook is called after each child-frame event and must remain a no-op + while the frame still has unfinished nodes. + """ + ... diff --git a/src/graphon/graph_engine/orchestration/dispatcher.py b/src/graphon/engine/dispatcher.py similarity index 58% rename from src/graphon/graph_engine/orchestration/dispatcher.py rename to src/graphon/engine/dispatcher.py index 9a7f1e8d..bb525b99 100644 --- a/src/graphon/graph_engine/orchestration/dispatcher.py +++ b/src/graphon/engine/dispatcher.py @@ -1,38 +1,34 @@ -"""Main dispatcher for processing events from workers.""" +"""Dispatch worker results onto the engine's state-transition thread.""" import logging import queue import threading from typing import final -from graphon.graph_engine.entities.tasks import ( - ContainerAwaitTask, - DispatchTask, -) -from graphon.graph_events.base import GraphNodeEventBase -from graphon.graph_events.node import ( +from graphon.engine_events.base import NodeEvent +from graphon.engine_events.node import ( NodeRunExceptionEvent, NodeRunFailedEvent, NodeRunModelPollingProgressEvent, NodeRunSucceededEvent, ) -from graphon.runtime.graph_runtime_state import GraphExecutionProtocol +from graphon.runtime.execution import GraphExecution -from ..command_processing import CommandProcessor -from ..event_management import EventManager -from ..event_management.event_handlers import EventHandler -from ..graph_state_manager import GraphStateManager -from ..worker_management import WorkerPool +from .command.processor import CommandProcessor +from .event.processor import NodeEventProcessor +from .event.stream import EventStream +from .scheduler import Scheduler +from .worker import ContainerAwaitTask, DispatchTask, WorkerPool logger = logging.getLogger(__name__) @final class Dispatcher: - """Main dispatcher that processes events from the event queue. + """Process worker dispatch tasks on one state-transition thread. - This runs in a separate thread and coordinates event processing - with timeout and completion detection. + Workers only execute nodes and enqueue results. The dispatcher serializes + those results, command processing, and execution completion detection. """ _COMMAND_TRIGGER_EVENTS = ( @@ -44,33 +40,33 @@ class Dispatcher: def __init__( self, - event_queue: queue.Queue[DispatchTask], - event_handler: EventHandler, - graph_execution: GraphExecutionProtocol, - state_manager: GraphStateManager, + dispatch_queue: queue.Queue[DispatchTask], + event_processor: NodeEventProcessor, + graph_execution: GraphExecution, + scheduler: Scheduler, command_processor: CommandProcessor, worker_pool: WorkerPool, - event_emitter: EventManager, + event_stream: EventStream, ) -> None: """Initialize the dispatcher. Args: - event_queue: Queue of events from workers - event_handler: Event handler registry for processing events + dispatch_queue: Queue of tasks produced by workers. + event_processor: Processor that applies node events to frame state. graph_execution: Aggregate tracking graph execution state - state_manager: Root frame execution state manager + scheduler: Root frame scheduler and completion tracker. command_processor: Processor for external engine commands worker_pool: Pool executing ready node tasks - event_emitter: Event manager to signal completion + event_stream: Stream to mark complete when dispatch ends. """ - self._event_queue = event_queue - self._event_handler = event_handler + self._dispatch_queue = dispatch_queue + self._event_processor = event_processor self._graph_execution = graph_execution - self._state_manager = state_manager + self._scheduler = scheduler self._command_processor = command_processor self._worker_pool = worker_pool - self._event_emitter = event_emitter + self._event_stream = event_stream self._thread: threading.Thread | None = None self._stop_event = threading.Event() @@ -83,7 +79,7 @@ def start(self) -> None: self._stop_event.clear() self._thread = threading.Thread( target=self._dispatcher_loop, - name="GraphDispatcher", + name="EngineDispatcher", daemon=True, ) self._thread.start() @@ -105,7 +101,7 @@ def _dispatcher_loop(self) -> None: finally: if not self._graph_execution.paused and not self._graph_execution.completed: self._graph_execution.complete() - self._event_emitter.mark_complete() + self._event_stream.mark_complete() def _run_until_exit(self) -> bool: self._process_commands() @@ -113,63 +109,62 @@ def _run_until_exit(self) -> bool: if ( self._graph_execution.aborted or self._graph_execution.error is not None - or self._state_manager.is_execution_complete() + or self._scheduler.is_execution_complete() ): return False if self._graph_execution.paused: - self._state_manager.defer_ready_tasks(self._worker_pool.drain()) + self._scheduler.defer_ready_tasks(self._worker_pool.drain()) return True - self._worker_pool.check_and_scale() self._dispatch_next_event() return False def _dispatch_next_event(self) -> None: try: - task = self._event_queue.get(timeout=0.1) + task = self._dispatch_queue.get(timeout=0.1) except queue.Empty: self._process_commands() return event = self._dispatch_task(task) - self._event_queue.task_done() + self._dispatch_queue.task_done() self._process_commands(event) def _drain_after_exit(self, paused: bool) -> None: self._process_commands() if paused: - self._drain_events_until_idle() - self._event_handler.snapshot_frames() + self._drain_dispatch_tasks_until_idle() + self._event_processor.snapshot_frames() else: - self._drain_event_queue() + self._drain_dispatch_queue() - def _process_commands(self, event: GraphNodeEventBase | None = None) -> None: + def _process_commands(self, event: NodeEvent | None = None) -> None: if event is None or isinstance(event, self._COMMAND_TRIGGER_EVENTS): self._command_processor.process_commands() - def _drain_event_queue(self) -> None: + def _drain_dispatch_queue(self) -> None: while True: try: - task = self._event_queue.get(block=False) + task = self._dispatch_queue.get(block=False) except queue.Empty: return self._dispatch_task(task) - self._event_queue.task_done() + self._dispatch_queue.task_done() - def _drain_events_until_idle(self) -> None: + def _drain_dispatch_tasks_until_idle(self) -> None: while not self._stop_event.is_set(): try: - task = self._event_queue.get(timeout=0.1) + task = self._dispatch_queue.get(timeout=0.1) except queue.Empty: if not self._worker_pool.has_current_tasks(): break continue event = self._dispatch_task(task) - self._event_queue.task_done() + self._dispatch_queue.task_done() self._process_commands(event) - self._drain_event_queue() + self._drain_dispatch_queue() - def _dispatch_task(self, task: DispatchTask) -> GraphNodeEventBase | None: + def _dispatch_task(self, task: DispatchTask) -> NodeEvent | None: if isinstance(task, ContainerAwaitTask): - self._event_handler.start_container(task) + self._event_processor.start_container(task) return None - self._event_handler.dispatch(task) + self._event_processor.dispatch(task) return task.event diff --git a/src/graphon/graph_engine/graph_engine.py b/src/graphon/engine/engine.py similarity index 58% rename from src/graphon/graph_engine/graph_engine.py rename to src/graphon/engine/engine.py index b91d5684..31299a99 100644 --- a/src/graphon/graph_engine/graph_engine.py +++ b/src/graphon/engine/engine.py @@ -1,8 +1,4 @@ -"""QueueBasedGraphEngine - Main orchestrator for queue-based workflow execution. - -This engine uses a modular architecture with separated packages following -Domain-Driven Design principles for improved maintainability and testability. -""" +"""Queue-based graph execution engine.""" from __future__ import annotations @@ -11,12 +7,10 @@ from collections.abc import Generator, Sequence from typing import final -from graphon.entities.workflow_start_reason import WorkflowStartReason -from graphon.graph.graph import Graph -from graphon.graph_events.base import ( - GraphEngineEvent, +from graphon.engine_events.base import ( + EngineEvent, ) -from graphon.graph_events.graph import ( +from graphon.engine_events.graph import ( GraphRunAbortedEvent, GraphRunFailedEvent, GraphRunPartialSucceededEvent, @@ -24,36 +18,32 @@ GraphRunStartedEvent, GraphRunSucceededEvent, ) -from graphon.runtime.graph_runtime_state import GraphRuntimeState +from graphon.entities.workflow_start_reason import WorkflowStartReason +from graphon.graph.graph import Graph +from graphon.runtime.execution import ROOT_FRAME_ID +from graphon.runtime.graph_runtime_state import RuntimeState from graphon.runtime.read_only_wrappers import ReadOnlyGraphRuntimeStateWrapper -from .command_channels import CommandChannel -from .command_processing import ( - AbortCommandHandler, - CommandProcessor, - PauseCommandHandler, - UpdateVariablesCommandHandler, +from .command.builtin.in_memory import InMemoryChannel +from .command.entities import AbortCommand +from .command.processor import CommandProcessor +from .command.protocol import CommandChannel +from .container_handler import ( + ContainerHandlerFactory, + IterationContainerHandler, + LoopContainerHandler, ) -from .config import GraphEngineConfig -from .container_handlers import ContainerHandlerFactory -from .entities.commands import AbortCommand, PauseCommand, UpdateVariablesCommand -from .entities.tasks import DispatchTask -from .error_handler import ErrorHandler -from .event_management import EventHandler, EventManager -from .frames import ExecutionFrame, FrameRegistry -from .graph_state_manager import GraphStateManager -from .graph_traversal import EdgeProcessor, SkipPropagator -from .iteration_container_handler import IterationContainerHandler -from .layers.base import GraphEngineLayer -from .loop_container_handler import LoopContainerHandler -from .orchestration import Dispatcher -from .ready_queue import ROOT_FRAME_ID, StartTask -from .worker_management import WorkerPool +from .dispatcher import Dispatcher +from .event.processor import NodeEventProcessor +from .event.stream import EventStream +from .frame import FrameRegistry +from .layer import Layer +from .ready_queue import StartTask +from .worker import DispatchTask, WorkerPool logger = logging.getLogger(__name__) -_DEFAULT_CONFIG = GraphEngineConfig() _DEFAULT_CONTAINER_HANDLER_FACTORIES: tuple[ContainerHandlerFactory, ...] = ( LoopContainerHandler, IterationContainerHandler, @@ -61,92 +51,73 @@ @final -class GraphEngine: - """Queue-based graph execution engine. - - Uses a modular architecture that delegates responsibilities to specialized - subsystems, following Domain-Driven Design and SOLID principles. - """ +class Engine: + """Coordinate graph scheduling, worker execution, and event delivery.""" def __init__( self, - workflow_id: str, graph: Graph, - graph_runtime_state: GraphRuntimeState, - command_channel: CommandChannel, - config: GraphEngineConfig = _DEFAULT_CONFIG, + graph_runtime_state: RuntimeState, + command_channel: CommandChannel | None = None, + workers: int = 5, container_handler_factories: Sequence[ContainerHandlerFactory] = (), ) -> None: - """Initialize the graph engine with all subsystems and dependencies.""" + """Build an engine for one graph execution. + + ``workers`` is the fixed number of node-execution threads owned by this + engine. It must be a positive integer; rejecting invalid values before + any runtime state is attached avoids creating an engine that can never + consume ready tasks. + + Args: + graph: Graph structure executed by this engine. + graph_runtime_state: Mutable runtime state associated with ``graph``. + command_channel: External command source for pause, abort, and updates. + A process-local in-memory channel is created when omitted. + workers: Fixed number of worker threads to create while running. + container_handler_factories: Additional container handler factories. + + Raises: + ValueError: If ``workers`` is not a positive integer. + + """ + if workers < 1: + msg = "workers must be a positive integer" + raise ValueError(msg) + # Bind runtime state to current workflow context self._graph = graph self._graph_runtime_state = graph_runtime_state - self._graph_runtime_state.attach_graph(graph) - self._command_channel = command_channel - self._layers: list[GraphEngineLayer] = [] + self._command_channel = ( + command_channel if command_channel is not None else InMemoryChannel() + ) + self._layers: list[Layer] = [] # Graph execution tracks the overall execution state self._graph_execution = self._graph_runtime_state.graph_execution - self._graph_execution.workflow_id = workflow_id - # Queue for events generated during execution - event_queue: queue.Queue[DispatchTask] = queue.Queue() + # Queue for state-transition work generated by workers. + dispatch_queue: queue.Queue[DispatchTask] = queue.Queue() # === State Management === - # Unified state manager handles all node state transitions and queue operations - self._state_manager = GraphStateManager( - self._graph, - self._graph_runtime_state, - ROOT_FRAME_ID, - ) self._frame_registry = FrameRegistry() - - # === Event Management === - # Event manager handles both collection and emission of events - self._event_manager = EventManager() - - # === Error Handling === - # Centralized error handler for graph execution errors - error_handler = ErrorHandler(self._graph, self._graph_execution) - - # === Graph Traversal Components === - # Propagates skip status through the graph when conditions aren't met - skip_propagator = SkipPropagator( + root_frame = self._frame_registry.create( + frame_id=ROOT_FRAME_ID, + container_id="", graph=self._graph, - state_manager=self._state_manager, + state=self._graph_runtime_state, ) + self._scheduler = root_frame.scheduler - # Processes edges to determine next nodes after execution - # Also handles conditional branching and route selection - edge_processor = EdgeProcessor( - graph=self._graph, - state_manager=self._state_manager, - skip_propagator=skip_propagator, - ) - self._frame_registry.register( - ExecutionFrame( - frame_id=ROOT_FRAME_ID, - graph=self._graph, - graph_runtime_state=self._graph_runtime_state, - state_manager=self._state_manager, - edge_processor=edge_processor, - error_handler=error_handler, - ), - ) + # === Event Streaming === + self._event_stream = EventStream(self._layers) # === Command Processing === # Processes external commands (e.g., abort requests) command_processor = CommandProcessor( command_channel=self._command_channel, graph_execution=self._graph_execution, - ) - - # Register command handlers - command_processor.register_handler(AbortCommand, AbortCommandHandler()) - command_processor.register_handler(PauseCommand, PauseCommandHandler()) - command_processor.register_handler( - UpdateVariablesCommand, - UpdateVariablesCommandHandler(self._graph_runtime_state.variable_pool), + variable_pool=self._graph_runtime_state.variable_pool, ) # === Worker Pool Setup === @@ -162,56 +133,59 @@ def __init__( # Create worker pool for parallel node execution self._worker_pool = WorkerPool( ready_queue=self._graph_runtime_state.ready_queue, - event_queue=event_queue, + dispatch_queue=dispatch_queue, frame_registry=self._frame_registry, layers=self._layers, execution_context=self._graph_runtime_state.execution_context, - config=config, + workers=workers, ) - # === Event Handler Registry === - # Central registry for handling all node execution events - event_handler = EventHandler( + # Applies node events to graph and frame state. + event_processor = NodeEventProcessor( graph_execution=self._graph_execution, - event_collector=self._event_manager, + event_stream=self._event_stream, frame_registry=self._frame_registry, container_handlers=self._container_handlers, ) # Dispatches events and manages execution flow self._dispatcher = Dispatcher( - event_queue=event_queue, - event_handler=event_handler, + dispatch_queue=dispatch_queue, + event_processor=event_processor, graph_execution=self._graph_execution, - state_manager=self._state_manager, + scheduler=self._scheduler, command_processor=command_processor, worker_pool=self._worker_pool, - event_emitter=self._event_manager, + event_stream=self._event_stream, ) # === Validation === - # Ensure all nodes share the same GraphRuntimeState instance + # Ensure all nodes share the same RuntimeState instance self._validate_graph_state_consistency() def _validate_graph_state_consistency(self) -> None: - """Validate that all nodes share the same GraphRuntimeState.""" + """Validate that all nodes share the same RuntimeState.""" expected_state_id = id(self._graph_runtime_state) for node in self._graph.nodes.values(): if id(node.graph_runtime_state) != expected_state_id: msg = ( - "GraphRuntimeState consistency violation: Node " + "RuntimeState consistency violation: Node " f"'{node.id}' has a different instance" ) raise ValueError(msg) - def layer(self, layer: GraphEngineLayer) -> GraphEngine: - """Add a layer for extending functionality.""" + def add_layer(self, layer: Layer) -> None: + """Register and bind one extension layer to this engine. + + The layer receives a read-only view of this engine's runtime state and + the command channel immediately. Its lifecycle hooks are invoked later + by :meth:`run`; registration itself does not start graph execution. + """ self._layers.append(layer) layer.initialize( ReadOnlyGraphRuntimeStateWrapper(self._graph_runtime_state), self._command_channel, ) - return self def request_abort(self, reason: str | None = None) -> None: """Queue an abort command for this engine.""" @@ -219,11 +193,11 @@ def request_abort(self, reason: str | None = None) -> None: AbortCommand(reason=reason or "User requested abort"), ) - def run(self) -> Generator[GraphEngineEvent, None, None]: + def run(self) -> Generator[EngineEvent, None, None]: """Execute the graph using the modular architecture. Yields: - `GraphEngineEvent` instances emitted during workflow execution. + `EngineEvent` instances emitted during workflow execution. """ try: @@ -233,14 +207,14 @@ def run(self) -> Generator[GraphEngineEvent, None, None]: error=str(error), exceptions_count=self._graph_execution.exceptions_count, ) - self._event_manager.notify_layers(failed_event) + self._event_stream.notify_layers(failed_event) yield failed_event raise finally: self._stop_execution() - def _run_graph(self) -> Generator[GraphEngineEvent, None, None]: - self._event_manager.reset() + def _run_graph(self) -> Generator[EngineEvent, None, None]: + self._event_stream.reset() self._initialize_layers() resume = self._graph_execution.started if resume: @@ -256,13 +230,13 @@ def _run_graph(self) -> Generator[GraphEngineEvent, None, None]: else WorkflowStartReason.INITIAL ), ) - self._event_manager.notify_layers(started_event) + self._event_stream.notify_layers(started_event) yield started_event self._start_execution(resume=resume) - yield from self._event_manager.emit_events() + yield from self._event_stream.emit_events() yield from self._emit_terminal_events() - def _emit_terminal_events(self) -> Generator[GraphEngineEvent, None, None]: + def _emit_terminal_events(self) -> Generator[EngineEvent, None, None]: if self._graph_execution.paused: pause_reasons = self._graph_execution.pause_reasons if not pause_reasons: @@ -273,7 +247,7 @@ def _emit_terminal_events(self) -> Generator[GraphEngineEvent, None, None]: reasons=pause_reasons, outputs=self._graph_runtime_state.outputs, ) - self._event_manager.notify_layers(paused_event) + self._event_stream.notify_layers(paused_event) yield paused_event return @@ -285,7 +259,7 @@ def _emit_terminal_events(self) -> Generator[GraphEngineEvent, None, None]: reason=abort_reason, outputs=self._graph_runtime_state.outputs, ) - self._event_manager.notify_layers(aborted_event) + self._event_stream.notify_layers(aborted_event) yield aborted_event return @@ -299,19 +273,18 @@ def _emit_terminal_events(self) -> Generator[GraphEngineEvent, None, None]: exceptions_count=exceptions_count, outputs=outputs, ) - self._event_manager.notify_layers(partial_event) + self._event_stream.notify_layers(partial_event) yield partial_event return succeeded_event = GraphRunSucceededEvent( outputs=outputs, ) - self._event_manager.notify_layers(succeeded_event) + self._event_stream.notify_layers(succeeded_event) yield succeeded_event def _initialize_layers(self) -> None: """Initialize layers with context.""" - self._event_manager.set_layers(self._layers) for layer in self._layers: try: layer.on_graph_start() @@ -329,25 +302,23 @@ def _start_execution(self, *, resume: bool) -> None: run_state = self._graph_runtime_state.get_container_run( frame_state.parent_invocation_id, ) - parent_node = self._frame_registry.get(run_state.frame_id).graph.nodes[ + parent_node = self._frame_registry[run_state.frame_id].graph.nodes[ run_state.node_id ] self._container_handlers[parent_node.node_type].restore_frame( frame_state, ) for run_state in self._graph_runtime_state.container_runs(): - self._frame_registry.get( - run_state.frame_id, - ).state_manager.track_unfinished(run_state.node_id) + self._frame_registry[run_state.frame_id].scheduler.track_unfinished( + run_state.node_id + ) ready_tasks = [ *self._graph_runtime_state.ready_queue.drain(), *self._graph_runtime_state.drain_deferred_ready_tasks(), ] for task in ready_tasks: if isinstance(task, StartTask): - self._frame_registry.get( - task.frame_id - ).state_manager.track_unfinished( + self._frame_registry[task.frame_id].scheduler.track_unfinished( task.node_id, ) @@ -359,7 +330,7 @@ def _start_execution(self, *, resume: bool) -> None: self._graph_runtime_state.enqueue_ready_task(task) else: root_node = self._graph.root_node - self._state_manager.enqueue_node(root_node.id) + self._scheduler.enqueue_node(root_node.id) self._dispatcher.start() @@ -386,6 +357,6 @@ def graph(self) -> Graph: return self._graph @property - def graph_runtime_state(self) -> GraphRuntimeState: + def graph_runtime_state(self) -> RuntimeState: """Get the graph runtime state.""" return self._graph_runtime_state diff --git a/src/graphon/engine/event/__init__.py b/src/graphon/engine/event/__init__.py new file mode 100644 index 00000000..9bcc5ad8 --- /dev/null +++ b/src/graphon/engine/event/__init__.py @@ -0,0 +1 @@ +"""Engine event streaming and node-event state transitions.""" diff --git a/src/graphon/graph_engine/error_handler.py b/src/graphon/engine/event/node_failure.py similarity index 66% rename from src/graphon/graph_engine/error_handler.py rename to src/graphon/engine/event/node_failure.py index dbd21e20..c4fab6f7 100644 --- a/src/graphon/graph_engine/error_handler.py +++ b/src/graphon/engine/event/node_failure.py @@ -1,44 +1,41 @@ -"""Main error handler that coordinates error strategies.""" +"""Node failure strategy selection.""" import logging import time from typing import assert_never, final -from graphon.enums import ( - ErrorStrategy as ErrorStrategyEnum, +from graphon.engine_events.base import NodeEvent +from graphon.engine_events.node import ( + NodeRunExceptionEvent, + NodeRunFailedEvent, + NodeRunRetryEvent, ) from graphon.enums import ( + ErrorStrategy, WorkflowNodeExecutionMetadataKey, WorkflowNodeExecutionStatus, ) from graphon.graph.graph import Graph -from graphon.graph_events.base import GraphNodeEventBase -from graphon.graph_events.node import ( - NodeRunExceptionEvent, - NodeRunFailedEvent, - NodeRunRetryEvent, -) from graphon.node_events.base import NodeRunResult -from graphon.runtime.graph_runtime_state import GraphExecutionProtocol +from graphon.runtime.execution import GraphExecution logger = logging.getLogger(__name__) @final -class ErrorHandler: - """Coordinates error handling strategies for node failures. +class NodeFailureHandler: + """Select the configured continuation for a failed node. - This acts as a facade for the various error strategies, - selecting and applying the appropriate strategy based on - node configuration. + A retry or configured error strategy becomes another node event. Returning + ``None`` tells the processor that the failure aborts its execution scope. """ def __init__( self, graph: Graph, - graph_execution: GraphExecutionProtocol, + graph_execution: GraphExecution, ) -> None: - """Initialize the error handler. + """Initialize failure handling for one frame's graph. Args: graph: The workflow graph @@ -48,22 +45,25 @@ def __init__( self._graph = graph self._graph_execution = graph_execution - def handle_node_failure( + def handle( self, *, frame_id: str, event: NodeRunFailedEvent, - ) -> GraphNodeEventBase | None: - """Handle a node failure event. + ) -> NodeEvent | None: + """Translate a failed node event into its configured continuation. - Selects and applies the appropriate error strategy based on - the node's configuration. + Retry eligibility is checked once here before the configured terminal + strategy is considered. The returned event is fed back into the node + event processor; ``None`` means no continuation exists and execution + should fail at the current scope. Args: - event: The node failure event + frame_id: Frame containing the failed node execution. + event: Failed node event to resolve. Returns: - Optional new event to process, or None to abort + Retry or exception event to process, or ``None`` to abort. """ node = self._graph.nodes[event.node_id] @@ -74,71 +74,49 @@ def handle_node_failure( ) retry_count = node_execution.retry_count - # First check if retry is configured and not exhausted if node.retry and retry_count < node.retry_config.max_retries: - result = self._handle_retry(event, retry_count) - if result: - # Retry count will be incremented when NodeRunRetryEvent is handled - return result + # Retry count is incremented when NodeRunRetryEvent is processed. + return self._handle_retry(event, retry_count) # Apply configured error strategy strategy = node.error_strategy match strategy: case None: - return self._handle_abort(event) - case ErrorStrategyEnum.FAIL_BRANCH: + logger.error( + "Node %s failed without a continuation strategy: %s", + event.node_id, + event.error, + ) + return None + case ErrorStrategy.FAIL_BRANCH: return self._handle_fail_branch(event) - case ErrorStrategyEnum.DEFAULT_VALUE: + case ErrorStrategy.DEFAULT_VALUE: return self._handle_default_value(event) case _: assert_never(strategy) - def _handle_abort(self, event: NodeRunFailedEvent) -> None: - """Handle error by aborting execution. - - This is the default strategy when no other strategy is specified. - It stops the entire graph execution when a node fails. - - Args: - event: The failure event - - """ - logger.error( - "Node %s failed with ABORT strategy: %s", - event.node_id, - event.error, - ) - # Return None to signal that execution should stop - def _handle_retry( self, event: NodeRunFailedEvent, retry_count: int, - ) -> NodeRunRetryEvent | None: + ) -> NodeRunRetryEvent: """Handle error by retrying the node. - This strategy re-attempts node execution up to a configured - maximum number of retries with configurable intervals. + Eligibility has already been established by :meth:`handle`; this helper + waits for the configured interval and builds the retry event. Args: event: The failure event retry_count: Current retry attempt count Returns: - NodeRunRetryEvent if retry should occur, None otherwise + Event requesting the next node attempt. """ node = self._graph.nodes[event.node_id] - # Check if we've exceeded max retries - if not node.retry or retry_count >= node.retry_config.max_retries: - return None - - # Wait for retry interval time.sleep(node.retry_config.retry_interval_seconds) - - # Create retry event return NodeRunRetryEvent( id=event.id, node_title=node.title, @@ -182,7 +160,7 @@ def _handle_fail_branch(self, event: NodeRunFailedEvent) -> NodeRunExceptionEven edge_source_handle="fail-branch", metadata={ WorkflowNodeExecutionMetadataKey.ERROR_STRATEGY: ( - ErrorStrategyEnum.FAIL_BRANCH + ErrorStrategy.FAIL_BRANCH ), }, ), @@ -223,7 +201,7 @@ def _handle_default_value(self, event: NodeRunFailedEvent) -> NodeRunExceptionEv outputs=outputs, metadata={ WorkflowNodeExecutionMetadataKey.ERROR_STRATEGY: ( - ErrorStrategyEnum.DEFAULT_VALUE + ErrorStrategy.DEFAULT_VALUE ), }, ), diff --git a/src/graphon/graph_engine/event_management/event_handlers.py b/src/graphon/engine/event/processor.py similarity index 74% rename from src/graphon/graph_engine/event_management/event_handlers.py rename to src/graphon/engine/event/processor.py index fcd1849a..c5f42f9f 100644 --- a/src/graphon/graph_engine/event_management/event_handlers.py +++ b/src/graphon/engine/event/processor.py @@ -1,32 +1,25 @@ -"""Event handler implementations for different event types.""" +"""Apply node events to execution state.""" import logging from collections.abc import Iterator, Mapping from functools import singledispatchmethod from typing import final -from graphon.enums import ( - ErrorStrategy, - NodeExecutionType, - NodeState, - NodeType, - WorkflowNodeExecutionStatus, -) -from graphon.graph_events.agent import NodeRunAgentLogEvent -from graphon.graph_events.base import GraphNodeEventBase -from graphon.graph_events.iteration import ( +from graphon.engine_events.agent import NodeRunAgentLogEvent +from graphon.engine_events.base import NodeEvent +from graphon.engine_events.iteration import ( NodeRunIterationFailedEvent, NodeRunIterationNextEvent, NodeRunIterationStartedEvent, NodeRunIterationSucceededEvent, ) -from graphon.graph_events.loop import ( +from graphon.engine_events.loop import ( NodeRunLoopFailedEvent, NodeRunLoopNextEvent, NodeRunLoopStartedEvent, NodeRunLoopSucceededEvent, ) -from graphon.graph_events.node import ( +from graphon.engine_events.node import ( NodeRunExceptionEvent, NodeRunFailedEvent, NodeRunModelPollingProgressEvent, @@ -39,51 +32,58 @@ NodeRunSucceededEvent, NodeRunVariableUpdatedEvent, ) +from graphon.enums import ( + ErrorStrategy, + NodeExecutionType, + NodeState, + NodeType, + WorkflowNodeExecutionStatus, +) from graphon.nodes.container_effects import ( ContainerExecutionResult, ContainerNodeRunResult, ) -from graphon.runtime.graph_runtime_state import GraphExecutionProtocol +from graphon.runtime.execution import ROOT_FRAME_ID, GraphExecution -from ..container_handlers import ContainerHandler -from ..entities.tasks import ContainerAwaitTask, TaskEvent -from ..frames import ExecutionFrame, FrameRegistry -from ..ready_queue import ROOT_FRAME_ID, ResumeTask, StartTask -from .event_manager import EventManager +from ..container_handler import ContainerHandler +from ..frame import ExecutionFrame, FrameRegistry +from ..ready_queue import ResumeTask, StartTask +from ..worker import ContainerAwaitTask, NodeEventTask +from .stream import EventStream logger = logging.getLogger(__name__) @final -class EventHandler: - """Registry of event handlers for different event types. +class NodeEventProcessor: + """Apply each node event type to its owning execution frame. - This centralizes the business logic for handling specific events, - keeping it separate from the routing and collection infrastructure. + This keeps node-event state transitions separate from worker dispatch and + external event streaming. """ def __init__( self, - graph_execution: GraphExecutionProtocol, - event_collector: EventManager, + graph_execution: GraphExecution, + event_stream: EventStream, frame_registry: FrameRegistry, container_handlers: Mapping[NodeType, ContainerHandler], ) -> None: - """Initialize the event handler registry. + """Initialize the node event processor. Args: graph_execution: Graph execution aggregate - event_collector: Event manager for collecting events + event_stream: Stream that collects processed engine events frame_registry: Registry of frame-local execution collaborators container_handlers: Engine-owned container handlers by node type """ self._graph_execution = graph_execution - self._event_collector = event_collector + self._event_stream = event_stream self._frame_registry = frame_registry self._container_handlers = container_handlers - def dispatch(self, task_event: TaskEvent) -> None: + def dispatch(self, task_event: NodeEventTask) -> None: """Handle any task-scoped node event. Args: @@ -91,21 +91,19 @@ def dispatch(self, task_event: TaskEvent) -> None: """ self._dispatch_event(frame_id=task_event.frame_id, event=task_event.event) - frame = self._frame_registry.get(task_event.frame_id) + frame = self._frame_registry[task_event.frame_id] handler = self._container_handler_for_frame(frame.frame_id) if handler is not None: - handler.complete_frame(frame) + handler.complete_frame_if_ready(frame) def start_container(self, task: ContainerAwaitTask) -> None: """Schedule child-frame work for a suspended container invocation.""" - root_runtime_state = self._frame_registry.get( - ROOT_FRAME_ID, - ).graph_runtime_state + root_runtime_state = self._frame_registry[ROOT_FRAME_ID].state run_state = root_runtime_state.get_container_run(task.invocation_id) - parent_frame = self._frame_registry.get(run_state.frame_id) + parent_frame = self._frame_registry[run_state.frame_id] node = parent_frame.graph.nodes[run_state.node_id] try: - self._container_handlers[node.node_type].start_await( + self._container_handlers[node.node_type].handle_request( invocation_id=task.invocation_id, request=task.request, ) @@ -131,11 +129,9 @@ def start_container(self, task: ContainerAwaitTask) -> None: def snapshot_frames(self) -> None: """Persist live child frames after workers have drained for a pause.""" - root_runtime_state = self._frame_registry.get( - ROOT_FRAME_ID, - ).graph_runtime_state + root_runtime_state = self._frame_registry[ROOT_FRAME_ID].state for frame_state in root_runtime_state.container_frames(): - frame = self._frame_registry.get(frame_state.frame_id) + frame = self._frame_registry[frame_state.frame_id] variable_pool_scope = ( "parent" if isinstance(frame_state.runtime_data.variable_pool, str) @@ -144,29 +140,30 @@ def snapshot_frames(self) -> None: root_runtime_state.put_container_frame( frame_state.model_copy( update={ - "runtime_data": frame.graph_runtime_state.snapshot_frame( + "runtime_data": frame.state.snapshot_frame( variable_pool_scope=variable_pool_scope, ), }, ), ) - def _dispatch_event(self, *, frame_id: str, event: GraphNodeEventBase) -> None: - frame = self._frame_registry.get(frame_id) + def _dispatch_event(self, *, frame_id: str, event: NodeEvent) -> None: + frame = self._frame_registry[frame_id] + event.container_id = frame.container_id for container_frame, handler in self._container_ancestors(frame_id): handler.prepare_frame_event(frame=container_frame, event=event) self._dispatch(event, frame=frame) @singledispatchmethod - def _dispatch(self, event: GraphNodeEventBase, *, frame: ExecutionFrame) -> None: + def _dispatch(self, event: NodeEvent, *, frame: ExecutionFrame) -> None: self._collect(frame=frame, event=event) logger.warning("Unhandled event type: %s", type(event).__name__) - def _collect(self, *, frame: ExecutionFrame, event: GraphNodeEventBase) -> None: + def _collect(self, *, frame: ExecutionFrame, event: NodeEvent) -> None: handler = self._container_handler_for_frame(frame.frame_id) - if handler is not None and not handler.should_collect(event=event): + if handler is not None and not handler.should_emit(event=event): return - self._event_collector.collect(event) + self._event_stream.collect(event) @_dispatch.register def _( @@ -205,7 +202,7 @@ def _(self, event: NodeRunStartedEvent, *, frame: ExecutionFrame) -> None: node_id=event.node_id, ) is_initial_attempt = node_execution.retry_count == 0 - frame.graph_runtime_state.increment_node_run_steps() + frame.state.increment_node_run_steps() # Collect the event only for the first attempt; retries remain silent if is_initial_attempt: @@ -218,7 +215,7 @@ def _(self, event: NodeRunVariableUpdatedEvent, *, frame: ExecutionFrame) -> Non The event is collected like other node events so parent/container engines can forward the updated payload to outer layers, including persistence listeners. """ - frame.graph_runtime_state.variable_pool.add( + frame.state.variable_pool.add( event.variable.selector, event.variable, ) @@ -237,17 +234,17 @@ def _(self, event: NodeRunSucceededEvent, *, frame: ExecutionFrame) -> None: def _(self, event: NodeRunPauseRequestedEvent, *, frame: ExecutionFrame) -> None: """Handle pause requests emitted by nodes.""" self._graph_execution.pause(event.reason) - frame.state_manager.finish_execution(event.node_id) + frame.scheduler.finish_execution(event.node_id) frame.graph.nodes[event.node_id].state = NodeState.UNKNOWN - frame.graph_runtime_state.defer_ready_task( + frame.state.defer_ready_task( StartTask(frame_id=frame.frame_id, node_id=event.node_id) ) - frame.state_manager.track_unfinished(event.node_id) + frame.scheduler.track_unfinished(event.node_id) self._collect(frame=frame, event=event) @_dispatch.register def _(self, event: NodeRunFailedEvent, *, frame: ExecutionFrame) -> None: - """Handle node failure using error handler. + """Resolve a node failure through its frame-local failure policy. Args: event: The node failed event @@ -256,9 +253,9 @@ def _(self, event: NodeRunFailedEvent, *, frame: ExecutionFrame) -> None: # Update domain model self._graph_execution.record_node_failure() - frame.graph_runtime_state.add_llm_usage(event.node_run_result.llm_usage) + frame.state.add_llm_usage(event.node_run_result.llm_usage) - result = frame.error_handler.handle_node_failure( + result = frame.failure_handler.handle( frame_id=frame.frame_id, event=event, ) @@ -273,7 +270,7 @@ def _(self, event: NodeRunFailedEvent, *, frame: ExecutionFrame) -> None: else: self._graph_execution.fail(RuntimeError(event.error)) self._collect(frame=frame, event=event) - frame.state_manager.finish_execution(event.node_id) + frame.scheduler.finish_execution(event.node_id) @_dispatch.register def _(self, event: NodeRunExceptionEvent, *, frame: ExecutionFrame) -> None: @@ -307,13 +304,13 @@ def _(self, event: NodeRunRetryEvent, *, frame: ExecutionFrame) -> None: node_execution.increment_retry() # Finish the previous attempt before re-queuing the node - frame.state_manager.finish_execution(event.node_id) + frame.scheduler.finish_execution(event.node_id) # Emit retry event for observers self._collect(frame=frame, event=event) # Re-queue node for execution - frame.state_manager.enqueue_node(event.node_id) + frame.scheduler.enqueue_node(event.node_id) def _complete_node( self, @@ -322,7 +319,7 @@ def _complete_node( event: NodeRunSucceededEvent | NodeRunExceptionEvent, follow_branch: bool, ) -> None: - frame.graph_runtime_state.add_llm_usage(event.node_run_result.llm_usage) + frame.state.add_llm_usage(event.node_run_result.llm_usage) self._store_node_outputs( frame=frame, node_id=event.node_id, @@ -330,25 +327,26 @@ def _complete_node( ) if follow_branch: - ready_nodes, edge_events = frame.edge_processor.handle_branch_completion( + ready_nodes, edge_events = frame.scheduler.handle_branch_completion( event.node_id, event.node_run_result.edge_source_handle, ) else: - ready_nodes, edge_events = frame.edge_processor.process_node_success( + ready_nodes, edge_events = frame.scheduler.process_node_success( event.node_id ) for edge_event in edge_events: - self._event_collector.collect(edge_event) + edge_event.container_id = frame.container_id + self._event_stream.collect(edge_event) for node_id in ready_nodes: - frame.state_manager.enqueue_node(node_id) + frame.scheduler.enqueue_node(node_id) node = frame.graph.nodes[event.node_id] if node.execution_type == NodeExecutionType.RESPONSE: - frame.graph_runtime_state.merge_response_outputs( + frame.state.merge_response_outputs( event.node_run_result.outputs, ) - frame.state_manager.finish_execution(event.node_id) + frame.scheduler.finish_execution(event.node_id) self._collect(frame=frame, event=event) def _store_node_outputs( @@ -366,7 +364,7 @@ def _store_node_outputs( """ for variable_name, variable_value in outputs.items(): - frame.graph_runtime_state.variable_pool.add( + frame.state.variable_pool.add( (node_id, variable_name), variable_value, ) @@ -381,18 +379,16 @@ def _container_ancestors( self, frame_id: str, ) -> Iterator[tuple[ExecutionFrame, ContainerHandler]]: - root_runtime_state = self._frame_registry.get( - ROOT_FRAME_ID, - ).graph_runtime_state + root_runtime_state = self._frame_registry[ROOT_FRAME_ID].state while frame_id != ROOT_FRAME_ID: frame_state = root_runtime_state.get_container_frame(frame_id) run_state = root_runtime_state.get_container_run( frame_state.parent_invocation_id, ) - parent_frame = self._frame_registry.get(run_state.frame_id) + parent_frame = self._frame_registry[run_state.frame_id] parent_node = parent_frame.graph.nodes[run_state.node_id] yield ( - self._frame_registry.get(frame_id), + self._frame_registry[frame_id], self._container_handlers[parent_node.node_type], ) frame_id = run_state.frame_id diff --git a/src/graphon/graph_engine/event_management/event_manager.py b/src/graphon/engine/event/stream.py similarity index 69% rename from src/graphon/graph_engine/event_management/event_manager.py rename to src/graphon/engine/event/stream.py index 6ea9a9b2..678953a1 100644 --- a/src/graphon/graph_engine/event_management/event_manager.py +++ b/src/graphon/engine/event/stream.py @@ -1,4 +1,4 @@ -"""Unified event manager for collecting and emitting events.""" +"""Thread-safe collection and delivery of engine events.""" import logging import threading @@ -7,9 +7,9 @@ from contextlib import contextmanager from typing import final -from graphon.graph_events.base import GraphEngineEvent +from graphon.engine_events.base import EngineEvent -from ..layers.base import GraphEngineLayer +from ..layer import Layer _logger = logging.getLogger(__name__) @@ -72,35 +72,48 @@ def write_lock(self) -> Generator: @final -class EventManager: - """Unified event manager that collects, buffers, and emits events. +class EventStream: + """Collect, buffer, and stream engine events. - This class combines event collection with event emission, providing - thread-safe event management with support for notifying layers and - streaming events to external consumers. + The stream is the single event boundary between the engine and external + consumers. It also notifies the engine's layers as events arrive. """ - def __init__(self) -> None: - """Initialize the event manager.""" - self._events: list[GraphEngineEvent] = [] - self._lock = ReadWriteLock() - self._layers: list[GraphEngineLayer] = [] - self._execution_complete = threading.Event() + def __init__(self, layers: list[Layer]) -> None: + """Initialize an event stream bound to the engine's live layer list. - def set_layers(self, layers: list[GraphEngineLayer]) -> None: - """Set the layers to notify on event collection. + The list is retained by reference so layers registered after engine + construction are visible to the stream without a second configuration + phase. Collected events are buffered until :meth:`emit_events` yields + them, while lifecycle events can notify the same layers without being + added to that buffer. Args: - layers: List of layers to notify + layers: Mutable list of layers owned by the engine. """ + self._events: list[EngineEvent] = [] + self._lock = ReadWriteLock() self._layers = layers + self._execution_complete = threading.Event() + + def notify_layers(self, event: EngineEvent) -> None: + """Notify all layers about an event without buffering it. + + Layer exceptions are caught and logged so one extension cannot disrupt + event delivery to the remaining layers or the engine itself. - def notify_layers(self, event: GraphEngineEvent) -> None: - """Notify registered layers about an event without buffering it.""" - self._notify_layers(event) + Args: + event: Event to send to every registered layer. + + """ + for layer in self._layers: + try: + layer.on_event(event) + except Exception: + _logger.exception("Error in layer on_event, layer_type=%s", type(layer)) - def collect(self, event: GraphEngineEvent) -> None: + def collect(self, event: EngineEvent) -> None: """Thread-safe method to collect an event. Args: @@ -110,14 +123,11 @@ def collect(self, event: GraphEngineEvent) -> None: with self._lock.write_lock(): self._events.append(event) - # NOTE: `_notify_layers` is intentionally called outside the critical section + # NOTE: `notify_layers` is intentionally called outside the critical section # to minimize lock contention and avoid blocking other readers or writers. - # - # The public `notify_layers` method also does not use a write lock, - # so protecting `_notify_layers` with a lock here is unnecessary. - self._notify_layers(event) + self.notify_layers(event) - def _get_new_events(self, start_index: int) -> list[GraphEngineEvent]: + def _get_new_events(self, start_index: int) -> list[EngineEvent]: """Get new events starting from a specific index. Args: @@ -150,11 +160,11 @@ def reset(self) -> None: self._events.clear() self._execution_complete.clear() - def emit_events(self) -> Generator[GraphEngineEvent, None, None]: + def emit_events(self) -> Generator[EngineEvent, None, None]: """Generator that yields events as they're collected. Yields: - GraphEngineEvent instances as they're processed + EngineEvent instances as they're processed """ yielded_count = 0 @@ -173,18 +183,3 @@ def emit_events(self) -> Generator[GraphEngineEvent, None, None]: # Small sleep to avoid busy waiting if not self._execution_complete.is_set() and not new_events: time.sleep(0.001) - - def _notify_layers(self, event: GraphEngineEvent) -> None: - """Notify all layers of an event. - - Layer exceptions are caught and logged to prevent disrupting collection. - - Args: - event: The event to send to layers - - """ - for layer in self._layers: - try: - layer.on_event(event) - except Exception: - _logger.exception("Error in layer on_event, layer_type=%s", type(layer)) diff --git a/src/graphon/engine/filter/__init__.py b/src/graphon/engine/filter/__init__.py new file mode 100644 index 00000000..9eec5f76 --- /dev/null +++ b/src/graphon/engine/filter/__init__.py @@ -0,0 +1,13 @@ +from .builtin.response_stream import ResponseStreamFilter +from .chain import filter_engine_events +from .protocol import ( + EngineEventFilter, + EngineEventFilterContext, +) + +__all__ = [ + "EngineEventFilter", + "EngineEventFilterContext", + "ResponseStreamFilter", + "filter_engine_events", +] diff --git a/src/graphon/engine/filter/builtin/__init__.py b/src/graphon/engine/filter/builtin/__init__.py new file mode 100644 index 00000000..078faa6d --- /dev/null +++ b/src/graphon/engine/filter/builtin/__init__.py @@ -0,0 +1 @@ +"""Implementation modules for event filters exported by the parent package.""" diff --git a/src/graphon/graph_engine/filters/response_stream.py b/src/graphon/engine/filter/builtin/response_stream.py similarity index 81% rename from src/graphon/graph_engine/filters/response_stream.py rename to src/graphon/engine/filter/builtin/response_stream.py index 5afe9f64..561e52d4 100644 --- a/src/graphon/graph_engine/filters/response_stream.py +++ b/src/graphon/engine/filter/builtin/response_stream.py @@ -8,34 +8,36 @@ from pydantic import BaseModel, Field -from graphon.enums import NodeExecutionType, NodeState -from graphon.graph_engine.filters.base import GraphEventFilterContext -from graphon.graph_events.base import GraphEngineEvent -from graphon.graph_events.graph import GraphRunStartedEvent -from graphon.graph_events.node import ( +from graphon.engine.filter.protocol import EngineEventFilterContext +from graphon.engine_events.base import EngineEvent +from graphon.engine_events.graph import GraphRunStartedEvent +from graphon.engine_events.node import ( NodeRunExceptionEvent, NodeRunReasoningChunkEvent, NodeRunStartedEvent, NodeRunStreamChunkEvent, NodeRunSucceededEvent, ) -from graphon.graph_events.traversal import GraphEdgeSkippedEvent, GraphEdgeTakenEvent +from graphon.engine_events.traversal import GraphEdgeSkippedEvent, GraphEdgeTakenEvent +from graphon.enums import NodeExecutionType, NodeState from graphon.nodes.base.template import Template, TextSegment, VariableSegment from graphon.runtime.graph_runtime_state import GraphProtocol, NodeProtocol from graphon.runtime.graph_runtime_state_protocol import ReadOnlyGraphRuntimeState -type NodeID = str -type EdgeID = str -type Selector = tuple[str, ...] +__all__ = ["ResponseStreamFilter"] + +type _NodeID = str +type _EdgeID = str +type _Selector = tuple[str, ...] @dataclass -class Path: +class _Path: """Blocking traversal edges that must be taken before a response can stream.""" - edges: list[EdgeID] = field(default_factory=list) + edges: list[_EdgeID] = field(default_factory=list) - def remove_edge(self, edge_id: EdgeID) -> None: + def remove_edge(self, edge_id: _EdgeID) -> None: if edge_id in self.edges: self.edges.remove(edge_id) @@ -44,7 +46,7 @@ def is_empty(self) -> bool: @dataclass -class ResponseSession: +class _ResponseSession: """Streaming cursor for one response node template.""" node_id: str @@ -52,12 +54,11 @@ class ResponseSession: index: int = 0 @classmethod - def from_node(cls, node: NodeProtocol) -> ResponseSession: + def from_node(cls, node: NodeProtocol) -> _ResponseSession: get_streaming_template = getattr(node, "get_streaming_template", None) if not callable(get_streaming_template): msg = ( - "ResponseSession.from_node requires " - "get_streaming_template() on response nodes" + "Response streaming requires get_streaming_template() on response nodes" ) raise TypeError(msg) return cls(node_id=node.id, template=get_streaming_template()) @@ -66,43 +67,43 @@ def is_complete(self) -> bool: return self.index >= len(self.template.segments) -class ResponseSessionState(BaseModel): +class _ResponseSessionState(BaseModel): """Serializable representation of a response session.""" node_id: str index: int = Field(default=0, ge=0) -class StreamBufferState(BaseModel): +class _StreamBufferState(BaseModel): """Serializable representation of buffered stream chunks.""" - selector: Selector + selector: _Selector events: list[NodeRunStreamChunkEvent] = Field(default_factory=list) -class StreamPositionState(BaseModel): +class _StreamPositionState(BaseModel): """Serializable representation for stream read positions.""" - selector: Selector + selector: _Selector position: int = Field(default=0, ge=0) @dataclass -class StreamBuffers: +class _StreamBuffers: """Buffered stream chunks plus per-selector read cursors.""" - events: dict[Selector, list[NodeRunStreamChunkEvent]] = field(default_factory=dict) - positions: dict[Selector, int] = field(default_factory=dict) - closed_selectors: set[Selector] = field(default_factory=set) + events: dict[_Selector, list[NodeRunStreamChunkEvent]] = field(default_factory=dict) + positions: dict[_Selector, int] = field(default_factory=dict) + closed_selectors: set[_Selector] = field(default_factory=set) @classmethod def from_state( cls, *, - buffers: Sequence[StreamBufferState], - positions: Sequence[StreamPositionState], - closed_selectors: Sequence[Selector], - ) -> StreamBuffers: + buffers: Sequence[_StreamBufferState], + positions: Sequence[_StreamPositionState], + closed_selectors: Sequence[_Selector], + ) -> _StreamBuffers: stream_buffers = cls( events={ tuple(buffer.selector): [ @@ -169,64 +170,62 @@ def close(self, selector: Sequence[str]) -> None: def is_closed(self, selector: Sequence[str]) -> bool: return tuple(selector) in self.closed_selectors - def dump_buffers(self) -> list[StreamBufferState]: + def dump_buffers(self) -> list[_StreamBufferState]: return [ - StreamBufferState( + _StreamBufferState( selector=selector, events=[event.model_copy(deep=True) for event in events], ) for selector, events in sorted(self.events.items()) ] - def dump_positions(self) -> list[StreamPositionState]: + def dump_positions(self) -> list[_StreamPositionState]: return [ - StreamPositionState(selector=selector, position=position) + _StreamPositionState(selector=selector, position=position) for selector, position in sorted(self.positions.items()) ] - def dump_closed_selectors(self) -> list[Selector]: + def dump_closed_selectors(self) -> list[_Selector]: return sorted(self.closed_selectors) -class ResponseStreamFilterState(BaseModel): +class _ResponseStreamFilterState(BaseModel): """Serialized snapshot of ResponseStreamFilter.""" type: Literal["ResponseStreamFilter"] = Field(default="ResponseStreamFilter") version: str = Field(default="1.0") response_nodes: Sequence[str] = Field(default_factory=list) - active_session: ResponseSessionState | None = None - waiting_sessions: Sequence[ResponseSessionState] = Field(default_factory=list) - pending_sessions: Sequence[ResponseSessionState] = Field(default_factory=list) + active_session: _ResponseSessionState | None = None + waiting_sessions: Sequence[_ResponseSessionState] = Field(default_factory=list) + pending_sessions: Sequence[_ResponseSessionState] = Field(default_factory=list) node_execution_ids: dict[str, str] = Field(default_factory=dict) paths_map: dict[str, list[list[str]]] = Field(default_factory=dict) - stream_buffers: Sequence[StreamBufferState] = Field(default_factory=list) - stream_positions: Sequence[StreamPositionState] = Field(default_factory=list) - closed_streams: Sequence[Selector] = Field(default_factory=list) + stream_buffers: Sequence[_StreamBufferState] = Field(default_factory=list) + stream_positions: Sequence[_StreamPositionState] = Field(default_factory=list) + closed_streams: Sequence[_Selector] = Field(default_factory=list) class ResponseStreamFilter: """Opt-in event filter that recreates legacy ordered response streaming.""" - filter_id = "graphon.response_stream.v1" - def __init__(self, *, pass_unmatched_chunks: bool = False) -> None: self._pass_unmatched_chunks = pass_unmatched_chunks self._graph: GraphProtocol | None = None self._runtime_state: ReadOnlyGraphRuntimeState | None = None - self._pending_state: ResponseStreamFilterState | None = None + self._pending_state: _ResponseStreamFilterState | None = None self._reset_run_state() def _reset_run_state(self) -> None: - self._active_session: ResponseSession | None = None - self._waiting_sessions: deque[ResponseSession] = deque() - self._stream_buffers = StreamBuffers() + self._active_session: _ResponseSession | None = None + self._waiting_sessions: deque[_ResponseSession] = deque() + self._stream_buffers = _StreamBuffers() self._response_nodes: set[str] = set() - self._paths_maps: dict[str, list[Path]] = {} + self._paths_maps: dict[str, list[_Path]] = {} self._node_execution_ids: dict[str, str] = {} - self._response_sessions: dict[str, ResponseSession] = {} - self._referenced_selectors: set[Selector] = set() + self._response_sessions: dict[str, _ResponseSession] = {} + self._referenced_selectors: set[_Selector] = set() - def initialize(self, context: GraphEventFilterContext) -> None: + def initialize(self, context: EngineEventFilterContext) -> None: pending_state = self._pending_state self._graph = cast(GraphProtocol, context.graph) self._runtime_state = context.runtime_state @@ -247,11 +246,11 @@ def initialize(self, context: GraphEventFilterContext) -> None: self._reset_run_state() raise - def on_event(self, event: GraphEngineEvent) -> Iterable[GraphEngineEvent]: + def on_event(self, event: EngineEvent) -> Iterable[EngineEvent]: self._ensure_initialized() match event: case GraphRunStartedEvent(): - output: Iterable[GraphEngineEvent] = [ + output: Iterable[EngineEvent] = [ event, *self._activate_initial_sessions(), ] @@ -272,7 +271,7 @@ def on_event(self, event: GraphEngineEvent) -> Iterable[GraphEngineEvent]: output = [event] return output - def flush(self) -> Iterable[GraphEngineEvent]: + def flush(self) -> Iterable[EngineEvent]: self._ensure_initialized() return self._try_flush() @@ -281,7 +280,7 @@ def dumps(self) -> str: return self._pending_state.model_dump_json() self._ensure_initialized() - state = ResponseStreamFilterState( + state = _ResponseStreamFilterState( response_nodes=sorted(self._response_nodes), active_session=self._serialize_session(self._active_session), waiting_sessions=[ @@ -314,8 +313,8 @@ def loads(self, data: str) -> None: self._apply_state(state) @staticmethod - def _parse_state(data: str) -> ResponseStreamFilterState: - state = ResponseStreamFilterState.model_validate_json(data) + def _parse_state(data: str) -> _ResponseStreamFilterState: + state = _ResponseStreamFilterState.model_validate_json(data) if state.type != "ResponseStreamFilter": msg = f"Invalid serialized data type: {state.type}" @@ -327,15 +326,15 @@ def _parse_state(data: str) -> ResponseStreamFilterState: return state - def _apply_state(self, state: ResponseStreamFilterState) -> None: + def _apply_state(self, state: _ResponseStreamFilterState) -> None: response_nodes = set(state.response_nodes) paths_maps = { - node_id: [Path(edges=list(path_edges)) for path_edges in paths] + node_id: [_Path(edges=list(path_edges)) for path_edges in paths] for node_id, paths in state.paths_map.items() } node_execution_ids = dict(state.node_execution_ids) - stream_buffers = StreamBuffers.from_state( + stream_buffers = _StreamBuffers.from_state( buffers=state.stream_buffers, positions=state.stream_positions, closed_selectors=state.closed_streams, @@ -354,7 +353,7 @@ def _apply_state(self, state: ResponseStreamFilterState) -> None: else None ) - referenced_selectors: set[Selector] = set() + referenced_selectors: set[_Selector] = set() for response_node_id in response_nodes: referenced_selectors.update( self._get_referenced_selectors(response_node_id) @@ -393,56 +392,56 @@ def _bound_runtime_state(self) -> ReadOnlyGraphRuntimeState: raise RuntimeError(msg) return self._runtime_state - def _register(self, response_node_id: NodeID) -> None: + def _register(self, response_node_id: _NodeID) -> None: if response_node_id in self._response_nodes: return self._response_nodes.add(response_node_id) self._paths_maps[response_node_id] = self._build_paths_map(response_node_id) response_node = self._bound_graph.nodes[response_node_id] - self._response_sessions[response_node_id] = ResponseSession.from_node( + self._response_sessions[response_node_id] = _ResponseSession.from_node( response_node, ) self._record_referenced_selectors(response_node_id) - def _record_referenced_selectors(self, response_node_id: NodeID) -> None: + def _record_referenced_selectors(self, response_node_id: _NodeID) -> None: self._referenced_selectors.update( self._get_referenced_selectors(response_node_id) ) def _get_referenced_selectors( self, - response_node_id: NodeID, - ) -> set[Selector]: + response_node_id: _NodeID, + ) -> set[_Selector]: response_node = self._bound_graph.nodes.get(response_node_id) if response_node is None: return set() - response_session = ResponseSession.from_node(response_node) + response_session = _ResponseSession.from_node(response_node) return { tuple(segment.selector) for segment in response_session.template.segments if isinstance(segment, VariableSegment) } - def _build_paths_map(self, response_node_id: NodeID) -> list[Path]: + def _build_paths_map(self, response_node_id: _NodeID) -> list[_Path]: root_node_id = self._bound_graph.root_node.id if root_node_id == response_node_id: - return [Path()] + return [_Path()] variable_selectors = self._get_response_variable_selectors(response_node_id) all_complete_paths = self._find_all_paths(root_node_id, response_node_id) return [ - Path(edges=self._get_blocking_edges(path, variable_selectors)) + _Path(edges=self._get_blocking_edges(path, variable_selectors)) for path in all_complete_paths ] def _get_response_variable_selectors( self, - response_node_id: NodeID, - ) -> set[Selector]: + response_node_id: _NodeID, + ) -> set[_Selector]: response_node = self._bound_graph.nodes[response_node_id] - response_session = ResponseSession.from_node(response_node) + response_session = _ResponseSession.from_node(response_node) return { tuple(segment.selector[:2]) for segment in response_session.template.segments @@ -451,18 +450,18 @@ def _get_response_variable_selectors( def _find_all_paths( self, - current_node_id: NodeID, - target_node_id: NodeID, - current_path: list[EdgeID] | None = None, - visited: set[NodeID] | None = None, - ) -> list[list[EdgeID]]: + current_node_id: _NodeID, + target_node_id: _NodeID, + current_path: list[_EdgeID] | None = None, + visited: set[_NodeID] | None = None, + ) -> list[list[_EdgeID]]: current_path = current_path or [] visited = visited or set() if current_node_id == target_node_id: return [current_path.copy()] next_visited = {current_node_id, *visited} - paths: list[list[EdgeID]] = [] + paths: list[list[_EdgeID]] = [] for edge in self._bound_graph.get_outgoing_edges(current_node_id): if edge.head in next_visited: continue @@ -478,9 +477,9 @@ def _find_all_paths( def _get_blocking_edges( self, - path: list[EdgeID], - variable_selectors: set[Selector], - ) -> list[EdgeID]: + path: list[_EdgeID], + variable_selectors: set[_Selector], + ) -> list[_EdgeID]: return [ edge_id for edge_id in path @@ -489,8 +488,8 @@ def _get_blocking_edges( def _is_blocking_edge( self, - edge_id: EdgeID, - variable_selectors: set[Selector], + edge_id: _EdgeID, + variable_selectors: set[_Selector], ) -> bool: edge = self._bound_graph.edges[edge_id] source_node = self._bound_graph.nodes[edge.tail] @@ -500,16 +499,16 @@ def _is_blocking_edge( NodeExecutionType.RESPONSE, )) or source_node.blocks_variable_output(variable_selectors) - def _activate_initial_sessions(self) -> list[GraphEngineEvent]: - events: list[GraphEngineEvent] = [] + def _activate_initial_sessions(self) -> list[EngineEvent]: + events: list[EngineEvent] = [] for response_node_id in sorted(self._response_nodes): paths = self._paths_maps.get(response_node_id, []) if any(path.is_empty() for path in paths): events.extend(self._active_or_queue_session(response_node_id)) return events - def _handle_edge_taken(self, edge_id: EdgeID) -> list[GraphEngineEvent]: - events: list[GraphEngineEvent] = [] + def _handle_edge_taken(self, edge_id: _EdgeID) -> list[EngineEvent]: + events: list[EngineEvent] = [] for response_node_id in sorted(self._response_nodes): paths = self._paths_maps.get(response_node_id) if paths is None: @@ -527,8 +526,8 @@ def _handle_edge_taken(self, edge_id: EdgeID) -> list[GraphEngineEvent]: def _active_or_queue_session( self, - node_id: NodeID, - ) -> list[GraphEngineEvent]: + node_id: _NodeID, + ) -> list[EngineEvent]: session = self._response_sessions.pop(node_id, None) if session is None: return [] @@ -543,7 +542,7 @@ def _active_or_queue_session( def _handle_stream_chunk( self, event: NodeRunStreamChunkEvent, - ) -> list[GraphEngineEvent]: + ) -> list[EngineEvent]: selector_key = tuple(event.selector) if selector_key in self._referenced_selectors: self._stream_buffers.append(event.selector, event) @@ -557,7 +556,7 @@ def _handle_stream_chunk( def _handle_reasoning_chunk( self, event: NodeRunReasoningChunkEvent, - ) -> list[GraphEngineEvent]: + ) -> list[EngineEvent]: if self._is_reasoning_visible(event): return [event] return [] @@ -573,7 +572,7 @@ def _is_reasoning_visible(self, event: NodeRunReasoningChunkEvent) -> bool: def _has_valid_reasoning_selector(event: NodeRunReasoningChunkEvent) -> bool: return tuple(event.selector) == (event.node_id, "reasoning_content") - def _has_runnable_reasoning_source(self, node_id: NodeID) -> bool: + def _has_runnable_reasoning_source(self, node_id: _NodeID) -> bool: source_node = self._bound_graph.nodes.get(node_id) return bool(source_node and source_node.state != NodeState.SKIPPED) @@ -591,8 +590,8 @@ def _is_reasoning_referenced_by_reached_session( for session in self._reached_sessions() ) - def _reached_sessions(self) -> tuple[ResponseSession, ...]: - sessions: list[ResponseSession] = [] + def _reached_sessions(self) -> tuple[_ResponseSession, ...]: + sessions: list[_ResponseSession] = [] if self._active_session is not None: sessions.append(self._active_session) sessions.extend(self._waiting_sessions) @@ -600,8 +599,8 @@ def _reached_sessions(self) -> tuple[ResponseSession, ...]: @staticmethod def _session_references_reasoning_source( - session: ResponseSession, - selector_prefixes: set[Selector], + session: _ResponseSession, + selector_prefixes: set[_Selector], ) -> bool: return any( isinstance(segment, VariableSegment) @@ -609,14 +608,14 @@ def _session_references_reasoning_source( for segment in session.template.segments ) - def _get_or_create_execution_id(self, node_id: NodeID) -> str: + def _get_or_create_execution_id(self, node_id: _NodeID) -> str: if node_id not in self._node_execution_ids: self._node_execution_ids[node_id] = str(uuid4()) return self._node_execution_ids[node_id] def _create_stream_chunk_event( self, - node_id: NodeID, + node_id: _NodeID, execution_id: str, selector: Sequence[str], chunk: str, @@ -672,6 +671,7 @@ def _process_variable_segment( id=execution_id, node_id=response_node.id, node_type=response_node.node_type, + container_id=event.container_id, selector=list(event.selector), chunk=event.chunk, is_final=event.is_final, @@ -727,7 +727,7 @@ def _process_text_segment( ) ] - def _get_text_segment_selector(self, response_node_id: NodeID) -> Sequence[str]: + def _get_text_segment_selector(self, response_node_id: _NodeID) -> Sequence[str]: response_node = self._bound_graph.nodes[response_node_id] get_streaming_text_selector = getattr( response_node, @@ -739,13 +739,13 @@ def _get_text_segment_selector(self, response_node_id: NodeID) -> Sequence[str]: return [str(part) for part in selector] return [response_node.id, "answer"] - def _try_flush(self) -> list[GraphEngineEvent]: + def _try_flush(self) -> list[EngineEvent]: if not self._active_session: return [] template = self._active_session.template response_node_id = self._active_session.node_id - events: list[GraphEngineEvent] = [] + events: list[EngineEvent] = [] while self._active_session.index < len(template.segments): segment = template.segments[self._active_session.index] @@ -774,7 +774,7 @@ def _try_flush(self) -> list[GraphEngineEvent]: return events - def _end_session(self, node_id: NodeID) -> list[GraphEngineEvent]: + def _end_session(self, node_id: _NodeID) -> list[EngineEvent]: if not self._active_session or self._active_session.node_id != node_id: return [] @@ -787,21 +787,21 @@ def _end_session(self, node_id: NodeID) -> list[GraphEngineEvent]: def _serialize_session( self, - session: ResponseSession | None, - ) -> ResponseSessionState | None: + session: _ResponseSession | None, + ) -> _ResponseSessionState | None: if session is None: return None - return ResponseSessionState(node_id=session.node_id, index=session.index) + return _ResponseSessionState(node_id=session.node_id, index=session.index) def _session_from_state( self, - session_state: ResponseSessionState, - ) -> ResponseSession: + session_state: _ResponseSessionState, + ) -> _ResponseSession: node = self._bound_graph.nodes.get(session_state.node_id) if node is None: msg = f"Unknown response node '{session_state.node_id}' in serialized state" raise ValueError(msg) - session = ResponseSession.from_node(node) + session = _ResponseSession.from_node(node) session.index = session_state.index return session diff --git a/src/graphon/graph_engine/filters/chain.py b/src/graphon/engine/filter/chain.py similarity index 63% rename from src/graphon/graph_engine/filters/chain.py rename to src/graphon/engine/filter/chain.py index 2c705581..05bbe112 100644 --- a/src/graphon/graph_engine/filters/chain.py +++ b/src/graphon/engine/filter/chain.py @@ -1,19 +1,19 @@ from collections.abc import Iterable -from graphon.graph_engine.filters.base import ( - GraphEventFilter, - GraphEventFilterContext, +from graphon.engine.filter.protocol import ( + EngineEventFilter, + EngineEventFilterContext, ) -from graphon.graph_events.base import GraphEngineEvent +from graphon.engine_events.base import EngineEvent -def filter_graph_events( - events: Iterable[GraphEngineEvent], +def filter_engine_events( + events: Iterable[EngineEvent], *, - context: GraphEventFilterContext, - filters: Iterable[GraphEventFilter], -) -> Iterable[GraphEngineEvent]: - """Apply graph event filters in registration order.""" + context: EngineEventFilterContext, + filters: Iterable[EngineEventFilter], +) -> Iterable[EngineEvent]: + """Apply engine event filters in registration order.""" filter_list = list(filters) for event_filter in filter_list: event_filter.initialize(context) @@ -35,14 +35,14 @@ def filter_graph_events( def _apply_filters( - event: GraphEngineEvent, + event: EngineEvent, *, - filters: list[GraphEventFilter], + filters: list[EngineEventFilter], start_index: int, -) -> Iterable[GraphEngineEvent]: +) -> Iterable[EngineEvent]: pending_events = [event] for event_filter in filters[start_index:]: - next_events: list[GraphEngineEvent] = [] + next_events: list[EngineEvent] = [] for pending_event in pending_events: next_events.extend(event_filter.on_event(pending_event)) pending_events = next_events diff --git a/src/graphon/engine/filter/protocol.py b/src/graphon/engine/filter/protocol.py new file mode 100644 index 00000000..722732cf --- /dev/null +++ b/src/graphon/engine/filter/protocol.py @@ -0,0 +1,40 @@ +from __future__ import annotations + +from collections.abc import Iterable +from dataclasses import dataclass +from typing import TYPE_CHECKING, Protocol + +from graphon.engine_events.base import EngineEvent +from graphon.graph.graph import Graph +from graphon.runtime.graph_runtime_state_protocol import ReadOnlyGraphRuntimeState +from graphon.runtime.read_only_wrappers import ReadOnlyGraphRuntimeStateWrapper + +if TYPE_CHECKING: + from graphon.engine.engine import Engine + + +@dataclass(frozen=True) +class EngineEventFilterContext: + """Run-scoped context available to engine event filters.""" + + graph: Graph + runtime_state: ReadOnlyGraphRuntimeState + + @classmethod + def from_engine(cls, engine: Engine) -> EngineEventFilterContext: + return cls( + graph=engine.graph, + runtime_state=ReadOnlyGraphRuntimeStateWrapper( + engine.graph_runtime_state, + ), + ) + + +class EngineEventFilter(Protocol): + """Event-to-event transform used outside Engine execution.""" + + def initialize(self, context: EngineEventFilterContext) -> None: ... + + def on_event(self, event: EngineEvent) -> Iterable[EngineEvent]: ... + + def flush(self) -> Iterable[EngineEvent]: ... diff --git a/src/graphon/engine/frame.py b/src/graphon/engine/frame.py new file mode 100644 index 00000000..aaaf5df5 --- /dev/null +++ b/src/graphon/engine/frame.py @@ -0,0 +1,278 @@ +"""Execution frame storage and construction.""" + +from __future__ import annotations + +from dataclasses import dataclass +from typing import Protocol, cast, final + +from graphon.graph.graph import Graph, NodeFactory +from graphon.runtime.container_state import FrameRuntimeData +from graphon.runtime.graph_runtime_state import RuntimeState +from graphon.runtime.variable_pool import VariablePool + +from .event.node_failure import NodeFailureHandler +from .scheduler import Scheduler + + +class RebindableNodeFactory(NodeFactory, Protocol): + def with_runtime_state( + self, + graph_runtime_state: RuntimeState, + ) -> RebindableNodeFactory: ... + + +@dataclass(frozen=True, slots=True) +class ExecutionFrame: + frame_id: str + graph: Graph + state: RuntimeState + scheduler: Scheduler + failure_handler: NodeFailureHandler + container_id: str = "" + + +@final +class FrameRegistry: + def __init__(self) -> None: + self._frames: dict[str, ExecutionFrame] = {} + + def register(self, frame: ExecutionFrame) -> None: + self._frames[frame.frame_id] = frame + + def __getitem__(self, frame_id: str) -> ExecutionFrame: + """Return a registered frame by its required identifier. + + Frame lookups are never optional during execution. Using subscription + syntax makes the existing ``KeyError`` behavior explicit instead of + resembling ``dict.get()``, which conventionally returns ``None``. + + Args: + frame_id: Identifier of the frame to retrieve. + + Returns: + The registered execution frame. + + """ + return self._frames[frame_id] + + def remove(self, frame_id: str) -> None: + del self._frames[frame_id] + + def create( + self, + *, + frame_id: str, + container_id: str, + graph: Graph, + state: RuntimeState, + ) -> ExecutionFrame: + """Create and register a fully wired frame. + + The runtime state owns graph attachment and any pending snapshot + restoration. This method only creates the frame-local scheduler and + failure handler after that attachment succeeds, so an invalid snapshot + can never leave a partially registered frame behind. + + Args: + frame_id: Unique frame ID used by scheduled tasks. + container_id: Direct owning container node ID; root uses ``""``. + graph: Graph structure executable by this frame. + state: Mutable runtime state owned by this graph execution. + + Returns: + The registered, fully wired execution frame. + + """ + state.attach_graph(graph) + scheduler = Scheduler( + graph, + state, + frame_id, + ) + frame = ExecutionFrame( + frame_id=frame_id, + graph=graph, + state=state, + scheduler=scheduler, + failure_handler=NodeFailureHandler(graph, state.graph_execution), + container_id=container_id, + ) + self.register(frame) + return frame + + def create_child( + self, + *, + frame_id: str, + parent_frame_id: str, + container_id: str, + root_node_id: str, + variable_pool: VariablePool, + ) -> ExecutionFrame: + """Create a child frame from its parent's container-scoped graph. + + The child receives an independent variable pool and runtime counters, + while queues and graph-wide execution state are inherited from the + parent. Only the graph config owned by ``container_id`` is materialized. + + Args: + frame_id: Unique ID for the new child frame. + parent_frame_id: Frame whose scoped config contains the child. + container_id: Container node that directly owns the child graph. + root_node_id: Entry node within the child graph. + variable_pool: Variable pool owned by the new child state. + + Returns: + The registered child execution frame. + + """ + parent_state = self[parent_frame_id].state + state = self._create_child_state( + parent_state=parent_state, + variable_pool=variable_pool, + ) + return self._create_child_with_state( + frame_id=frame_id, + parent_frame_id=parent_frame_id, + container_id=container_id, + root_node_id=root_node_id, + state=state, + ) + + def restore_child( + self, + *, + frame_id: str, + parent_frame_id: str, + container_id: str, + root_node_id: str, + runtime_data: FrameRuntimeData, + variable_pool: VariablePool, + ) -> ExecutionFrame: + """Restore a child frame from container-agnostic runtime data. + + Container handlers extract their own frame-specific fields before + calling this method. The registry therefore needs no knowledge of Loop, + Iteration, or downstream container state models. Runtime counters and + saved graph states are restored before the graph is attached. + + Args: + frame_id: Identifier of the child frame being restored. + parent_frame_id: Frame whose scoped config contains the child. + container_id: Container node that directly owns the child graph. + root_node_id: Entry node within the restored child graph. + runtime_data: Generic persisted data for the child runtime. + variable_pool: Resolved variable pool owned by the child. + + Returns: + The registered and restored child execution frame. + + """ + parent_state = self[parent_frame_id].state + state = self._create_child_state( + parent_state=parent_state, + variable_pool=variable_pool, + runtime_data=runtime_data, + ) + return self._create_child_with_state( + frame_id=frame_id, + parent_frame_id=parent_frame_id, + container_id=container_id, + root_node_id=root_node_id, + state=state, + ) + + def _create_child_with_state( + self, + *, + frame_id: str, + parent_frame_id: str, + container_id: str, + root_node_id: str, + state: RuntimeState, + ) -> ExecutionFrame: + """Build a scoped child graph around an already prepared state. + + The parent's retained graph config and node factory are the only graph + construction inputs. Rebinding the factory prevents child nodes from + observing the parent's mutable runtime state. + + Args: + frame_id: Unique ID for the new child frame. + parent_frame_id: Frame whose graph owns the container config. + container_id: Container node that scopes the child graph. + root_node_id: Entry node within the child graph. + state: Prepared runtime state to bind to child nodes. + + Returns: + The registered child execution frame. + + Raises: + RuntimeError: If the parent lacks graph config or a node factory. + + """ + parent_graph = self[parent_frame_id].graph + graph_config = parent_graph.graph_config + if graph_config is None: + msg = "Parent graph does not carry graph_config for frame creation." + raise RuntimeError(msg) + node_factory = parent_graph.node_factory + if node_factory is None: + msg = "Parent graph does not carry node_factory for frame creation." + raise RuntimeError(msg) + + rebound_factory = cast(RebindableNodeFactory, node_factory).with_runtime_state( + state, + ) + graph = Graph.init( + graph_config=graph_config, + node_factory=rebound_factory, + root_node_id=root_node_id, + container_id=container_id, + ) + return self.create( + frame_id=frame_id, + container_id=container_id, + graph=graph, + state=state, + ) + + @staticmethod + def _create_child_state( + *, + parent_state: RuntimeState, + variable_pool: VariablePool, + runtime_data: FrameRuntimeData | None = None, + ) -> RuntimeState: + """Create a child runtime state with parent-owned shared services. + + Fresh children start with empty counters and outputs. Restored children + receive their saved usage, outputs, steps, and pending graph states. + Both paths reuse the parent's ready queues and graph-wide execution so + scheduling and terminal status remain coordinated across frames. + + Args: + parent_state: Runtime state of the frame owning the child container. + variable_pool: Independent variable pool for the child frame. + runtime_data: Optional persisted child data to restore. + + Returns: + An unattached runtime state ready for child graph construction. + + """ + state = RuntimeState( + variable_pool=variable_pool, + start_at=parent_state.start_at, + llm_usage=None if runtime_data is None else runtime_data.llm_usage, + outputs=None if runtime_data is None else dict(runtime_data.outputs), + node_run_steps=0 if runtime_data is None else runtime_data.node_run_steps, + ready_queue=parent_state.ready_queue, + deferred_ready_queue=parent_state.deferred_ready_queue, + graph_execution=parent_state.graph_execution, + ) + if runtime_data is not None: + state.restore_graph_state( + node_states=runtime_data.graph_node_states, + edge_states=runtime_data.graph_edge_states, + ) + return state diff --git a/src/graphon/engine/layer/README.md b/src/graphon/engine/layer/README.md new file mode 100644 index 00000000..0ed6bdaa --- /dev/null +++ b/src/graphon/engine/layer/README.md @@ -0,0 +1,40 @@ +# Layers + +Pluggable middleware for engine extensions. + +## Components + +### Layer (base) + +Base class with optional lifecycle hooks for layers. + +- `initialize()` - Receive runtime context (runtime state is bound here and always available to hooks) +- `on_graph_start()` - Execution start hook +- `on_event()` - Process all events +- `on_graph_end()` - Execution end hook + +## Usage + +```python +from graphon.engine.layer import Layer +from graphon.engine_events import EngineEvent, NodeRunSucceededEvent + + +class MetricsLayer(Layer): + def __init__(self): + """Create storage for elapsed time collected during one engine run.""" + super().__init__() + self.metrics: dict[str, float] = {} + + def on_graph_start(self) -> None: + """Reset collected metrics before the engine starts a new graph run.""" + self.metrics.clear() + + def on_event(self, event: EngineEvent) -> None: + """Record elapsed time when a node run succeeds.""" + if isinstance(event, NodeRunSucceededEvent): + self.metrics[event.node_id] = event.node_run_result.elapsed_time +``` + +`engine.add_layer()` binds the read-only runtime state before execution, so +`graph_runtime_state` is always available inside layer hooks. diff --git a/src/graphon/engine/layer/__init__.py b/src/graphon/engine/layer/__init__.py new file mode 100644 index 00000000..ac6d56da --- /dev/null +++ b/src/graphon/engine/layer/__init__.py @@ -0,0 +1,13 @@ +"""Layer system for Engine extensibility. + +This module provides the layer infrastructure for extending Engine functionality +with middleware-like components that can observe events and interact with execution. +""" + +from .base import Layer +from .builtin.execution_limits import ExecutionLimitsLayer + +__all__ = [ + "ExecutionLimitsLayer", + "Layer", +] diff --git a/src/graphon/graph_engine/layers/base.py b/src/graphon/engine/layer/base.py similarity index 70% rename from src/graphon/graph_engine/layers/base.py rename to src/graphon/engine/layer/base.py index 98884195..003d654e 100644 --- a/src/graphon/graph_engine/layers/base.py +++ b/src/graphon/engine/layer/base.py @@ -1,41 +1,28 @@ -"""Base layer class for GraphEngine extensions. +"""Base class for Engine extensions. -This module provides the abstract base class for implementing layers that can -intercept and respond to GraphEngine events. +This module defines the lifecycle hooks and shared runtime binding used by layers +that intercept and respond to Engine events. """ -from abc import ABC, abstractmethod - -from graphon.graph_engine.command_channels.protocol import CommandChannel -from graphon.graph_events.base import ( - GraphEngineEvent, - GraphNodeEventBase, +from graphon.engine.command.protocol import CommandChannel +from graphon.engine_events.base import ( + EngineEvent, + NodeEvent, ) from graphon.nodes.base.node import Node from graphon.runtime.graph_runtime_state_protocol import ReadOnlyGraphRuntimeState -class GraphEngineLayerNotInitializedError(Exception): - """Raised when a layer's runtime state is accessed before initialization.""" - - def __init__(self, layer_name: str | None = None) -> None: - name = layer_name or "GraphEngineLayer" - super().__init__( - f"{name} runtime state is not initialized. " - "Bind the layer to a GraphEngine before access.", - ) - - -class GraphEngineLayer(ABC): - """Abstract base class for GraphEngine layers. +class Layer: + """Base class for Engine layers. Layers are middleware-like components that can: - - Observe all events emitted by the GraphEngine + - Observe all events emitted by the Engine - Access the graph runtime state - Send commands to control execution - Subclasses should override the constructor to accept configuration parameters, - then implement the three lifecycle methods. + Subclasses override only the lifecycle hooks they need. The default hooks + are no-ops, so event-only and node-only layers do not need placeholder methods. """ def __init__(self) -> None: @@ -46,7 +33,11 @@ def __init__(self) -> None: @property def graph_runtime_state(self) -> ReadOnlyGraphRuntimeState: if self._graph_runtime_state is None: - raise GraphEngineLayerNotInitializedError(type(self).__name__) + msg = ( + f"{type(self).__name__} runtime state is not initialized. " + "Bind the layer to an Engine before access." + ) + raise RuntimeError(msg) return self._graph_runtime_state def initialize( @@ -56,8 +47,8 @@ def initialize( ) -> None: """Initialize the layer with engine dependencies. - Called by GraphEngine to inject the read-only runtime state and command channel. - This is invoked when the layer is registered with a `GraphEngine` instance. + Called by Engine to inject the read-only runtime state and command channel. + This is invoked when the layer is registered with an `Engine` instance. Implementations should be idempotent. Args: @@ -68,7 +59,6 @@ def initialize( self._graph_runtime_state = graph_runtime_state self.command_channel = command_channel - @abstractmethod def on_graph_start(self) -> None: """Called when graph execution starts. @@ -76,8 +66,7 @@ def on_graph_start(self) -> None: are executed. Layers can use this to set up resources or log start information. """ - @abstractmethod - def on_event(self, event: GraphEngineEvent) -> None: + def on_event(self, event: EngineEvent) -> None: """Called for every event emitted by the engine. This method receives all events generated during graph execution, including: @@ -91,7 +80,6 @@ def on_event(self, event: GraphEngineEvent) -> None: """ - @abstractmethod def on_graph_end(self, error: Exception | None) -> None: """Called when graph execution ends. @@ -121,7 +109,7 @@ def on_node_run_end( self, node: Node, error: Exception | None, - result_event: GraphNodeEventBase | None = None, + result_event: NodeEvent | None = None, ) -> None: """Called after a node finishes execution. diff --git a/src/graphon/engine/layer/builtin/__init__.py b/src/graphon/engine/layer/builtin/__init__.py new file mode 100644 index 00000000..312db475 --- /dev/null +++ b/src/graphon/engine/layer/builtin/__init__.py @@ -0,0 +1 @@ +"""Implementation modules for layers exported by the parent package.""" diff --git a/src/graphon/graph_engine/layers/execution_limits.py b/src/graphon/engine/layer/builtin/execution_limits.py similarity index 64% rename from src/graphon/graph_engine/layers/execution_limits.py rename to src/graphon/engine/layer/builtin/execution_limits.py index 3b0078ef..b18c1d0d 100644 --- a/src/graphon/graph_engine/layers/execution_limits.py +++ b/src/graphon/engine/layer/builtin/execution_limits.py @@ -1,4 +1,4 @@ -"""Execution limits layer for GraphEngine. +"""Execution limits layer for Engine. This layer monitors workflow execution to enforce limits on: - Maximum execution steps @@ -9,28 +9,20 @@ import logging import time -from enum import StrEnum -from typing import assert_never, final, override +from typing import final, override -from graphon.graph_engine.entities.commands import AbortCommand, CommandType -from graphon.graph_engine.layers.base import GraphEngineLayer -from graphon.graph_events.base import GraphEngineEvent -from graphon.graph_events.node import ( +from graphon.engine.command.entities import AbortCommand +from graphon.engine.layer.base import Layer +from graphon.engine_events.base import EngineEvent +from graphon.engine_events.node import ( NodeRunFailedEvent, NodeRunStartedEvent, NodeRunSucceededEvent, ) -class LimitType(StrEnum): - """Types of execution limits that can be exceeded.""" - - STEP_LIMIT = "step_limit" - TIME_LIMIT = "time_limit" - - @final -class ExecutionLimitsLayer(GraphEngineLayer): +class ExecutionLimitsLayer(Layer): """Layer that enforces execution limits for workflows. Monitors: @@ -74,7 +66,7 @@ def on_graph_start(self) -> None: self.logger.debug("Execution limits monitoring started") @override - def on_event(self, event: GraphEngineEvent) -> None: + def on_event(self, event: EngineEvent) -> None: """Called for every event emitted by the engine. Monitors execution progress and enforces limits. @@ -87,11 +79,22 @@ def on_event(self, event: GraphEngineEvent) -> None: self.step_count += 1 self.logger.debug("Step %d started: %s", self.step_count, event.node_id) case NodeRunSucceededEvent() | NodeRunFailedEvent(): - if self._reached_step_limitation(): - self._send_abort_command(LimitType.STEP_LIMIT) - - if self._reached_time_limitation(): - self._send_abort_command(LimitType.TIME_LIMIT) + if self._step_limit_exceeded(): + reason = ( + "Maximum execution steps exceeded: " + f"{self.step_count} > {self.max_steps}" + ) + elif ( + start_time := self.start_time + ) is not None and self._time_limit_exceeded(): + elapsed_time = time.time() - start_time + reason = ( + "Maximum execution time exceeded: " + f"{elapsed_time:.2f}s > {self.max_time}s" + ) + else: + return + self._send_abort_command(reason) case _: pass @@ -109,22 +112,22 @@ def on_graph_end(self, error: Exception | None) -> None: total_time, ) - def _reached_step_limitation(self) -> bool: + def _step_limit_exceeded(self) -> bool: """Check if step count limit has been exceeded.""" return self.step_count > self.max_steps - def _reached_time_limitation(self) -> bool: + def _time_limit_exceeded(self) -> bool: """Check if time limit has been exceeded.""" return ( self.start_time is not None and (time.time() - self.start_time) > self.max_time ) - def _send_abort_command(self, limit_type: LimitType) -> None: + def _send_abort_command(self, reason: str) -> None: """Send abort command due to limit violation. Args: - limit_type: Type of limit exceeded + reason: Human-readable description of the exceeded limit. """ if ( @@ -135,13 +138,11 @@ def _send_abort_command(self, limit_type: LimitType) -> None: ): return - reason = self._build_abort_reason(limit_type) - self.logger.warning("Execution limit exceeded: %s", reason) try: # Send abort command to the engine - abort_command = AbortCommand(command_type=CommandType.ABORT, reason=reason) + abort_command = AbortCommand(reason=reason) self.command_channel.send_command(abort_command) # Mark that abort has been sent to prevent duplicate commands @@ -151,23 +152,3 @@ def _send_abort_command(self, limit_type: LimitType) -> None: except Exception: self.logger.exception("Failed to send abort command") - - def send_abort_command(self, limit_type: LimitType) -> None: - """Send an abort command when tests or callers need explicit control.""" - self._send_abort_command(limit_type) - - def _build_abort_reason(self, limit_type: LimitType) -> str: - match limit_type: - case LimitType.STEP_LIMIT: - return ( - f"Maximum execution steps exceeded: " - f"{self.step_count} > {self.max_steps}" - ) - case LimitType.TIME_LIMIT: - elapsed_time = time.time() - self.start_time if self.start_time else 0 - return ( - f"Maximum execution time exceeded: " - f"{elapsed_time:.2f}s > {self.max_time}s" - ) - case _: - assert_never(limit_type) diff --git a/src/graphon/graph_engine/ready_queue/__init__.py b/src/graphon/engine/ready_queue/__init__.py similarity index 64% rename from src/graphon/graph_engine/ready_queue/__init__.py rename to src/graphon/engine/ready_queue/__init__.py index f0d4e69e..058f3866 100644 --- a/src/graphon/graph_engine/ready_queue/__init__.py +++ b/src/graphon/engine/ready_queue/__init__.py @@ -1,15 +1,13 @@ -"""Ready queue implementations and serialized state helpers for GraphEngine.""" +"""Ready queue implementations and serialized state helpers for Engine.""" from graphon.runtime.ready_queue import ReadyQueue +from .entities import ReadyTask, ResumeTask, StartTask from .in_memory import InMemoryReadyQueue -from .protocol import ROOT_FRAME_ID, ReadyQueueState, ReadyTask, ResumeTask, StartTask __all__ = [ - "ROOT_FRAME_ID", "InMemoryReadyQueue", "ReadyQueue", - "ReadyQueueState", "ReadyTask", "ResumeTask", "StartTask", diff --git a/src/graphon/graph_engine/ready_queue/protocol.py b/src/graphon/engine/ready_queue/entities.py similarity index 50% rename from src/graphon/graph_engine/ready_queue/protocol.py rename to src/graphon/engine/ready_queue/entities.py index 307a639c..c26d54e1 100644 --- a/src/graphon/graph_engine/ready_queue/protocol.py +++ b/src/graphon/engine/ready_queue/entities.py @@ -1,13 +1,11 @@ -"""Serialized state models for GraphEngine ready queue implementations.""" +"""Tasks consumed by Engine ready queue implementations.""" -from typing import Annotated, Final, Literal +from typing import Annotated, Literal from pydantic import BaseModel, ConfigDict, Field from graphon.nodes.container_effects import ContainerRunResult -ROOT_FRAME_ID: Final = "root" - class StartTask(BaseModel): """Task that starts a node invocation inside an execution frame.""" @@ -30,22 +28,3 @@ class ResumeTask(BaseModel): ReadyTask = Annotated[StartTask | ResumeTask, Field(discriminator="kind")] - - -class ReadyQueueStateV1(BaseModel): - """Ready queue state produced before frame-aware tasks.""" - - type: Literal["InMemoryReadyQueue"] - version: Literal["1.0"] - items: tuple[str, ...] - - -class ReadyQueueState(BaseModel): - """Pydantic model for serialized ready queue state. - - This defines the structure of the data returned by dumps() - and expected by loads() for ready queue serialization. - """ - - version: Literal["2.0"] - items: tuple[ReadyTask, ...] diff --git a/src/graphon/graph_engine/ready_queue/in_memory.py b/src/graphon/engine/ready_queue/in_memory.py similarity index 85% rename from src/graphon/graph_engine/ready_queue/in_memory.py rename to src/graphon/engine/ready_queue/in_memory.py index 0f0d5487..3828ab95 100644 --- a/src/graphon/graph_engine/ready_queue/in_memory.py +++ b/src/graphon/engine/ready_queue/in_memory.py @@ -5,21 +5,36 @@ """ import queue -from typing import Annotated, final +from typing import Annotated, Literal, final -from pydantic import Field, TypeAdapter +from pydantic import BaseModel, Field, TypeAdapter -from .protocol import ( - ROOT_FRAME_ID, - ReadyQueueState, - ReadyQueueStateV1, +from graphon.runtime.execution import ROOT_FRAME_ID + +from .entities import ( ReadyTask, StartTask, ) + +class _ReadyQueueStateV1(BaseModel): + """Ready queue state produced before frame-aware tasks.""" + + type: Literal["InMemoryReadyQueue"] + version: Literal["1.0"] + items: tuple[str, ...] + + +class _ReadyQueueState(BaseModel): + """Validated serialized state for the in-memory ready queue.""" + + version: Literal["2.0"] + items: tuple[ReadyTask, ...] + + _READY_QUEUE_STATE_ADAPTER = TypeAdapter( Annotated[ - ReadyQueueStateV1 | ReadyQueueState, + _ReadyQueueStateV1 | _ReadyQueueState, Field(discriminator="version"), ], ) @@ -108,7 +123,7 @@ def dumps(self) -> str: # callers must quiesce producers and consumers while serializing. items = self.drain() try: - state = ReadyQueueState( + state = _ReadyQueueState( version="2.0", items=tuple(items), ) @@ -125,7 +140,7 @@ def loads(self, data: str) -> None: """ state = _READY_QUEUE_STATE_ADAPTER.validate_json(data) - if isinstance(state, ReadyQueueStateV1): + if isinstance(state, _ReadyQueueStateV1): items: tuple[ReadyTask, ...] = tuple( StartTask(frame_id=ROOT_FRAME_ID, node_id=node_id) for node_id in state.items diff --git a/src/graphon/engine/scheduler.py b/src/graphon/engine/scheduler.py new file mode 100644 index 00000000..ee3db5db --- /dev/null +++ b/src/graphon/engine/scheduler.py @@ -0,0 +1,402 @@ +"""Frame-local graph scheduling, traversal, and execution tracking.""" + +from collections.abc import Sequence +from typing import TypedDict, final + +from graphon.engine_events.traversal import GraphEdgeSkippedEvent, GraphEdgeTakenEvent +from graphon.enums import NodeExecutionType, NodeState +from graphon.graph.edge import Edge +from graphon.graph.graph import Graph +from graphon.runtime.graph_runtime_state import RuntimeState + +from .ready_queue import ReadyTask, StartTask + +type GraphTraversalEvent = GraphEdgeTakenEvent | GraphEdgeSkippedEvent + + +class _EdgeStateAnalysis(TypedDict): + """Analysis result for edge states.""" + + has_unknown: bool + has_taken: bool + all_skipped: bool + + +@final +class Scheduler: + def __init__( + self, + graph: Graph, + state: RuntimeState, + frame_id: str, + ) -> None: + """Initialize frame-local scheduling and traversal state. + + Args: + graph: The workflow graph + state: Runtime state owning ready task queues + frame_id: Execution frame managed by this instance + + """ + self._graph = graph + self._state = state + self._frame_id = frame_id + self._unfinished_nodes: set[str] = set() + + # ============= Node State Operations ============= + + def enqueue_node(self, node_id: str) -> None: + """Mark a node as TAKEN and add its task to the ready queue. + + This combines the state transition and enqueueing operations + that always occur together when preparing a node for execution. + + Args: + node_id: The ID of the node to enqueue + + """ + self._graph.nodes[node_id].state = NodeState.TAKEN + self._unfinished_nodes.add(node_id) + self._state.enqueue_ready_task( + StartTask(frame_id=self._frame_id, node_id=node_id), + ) + + def mark_node_skipped(self, node_id: str) -> None: + """Mark a node as SKIPPED. + + Args: + node_id: The ID of the node to skip + + """ + self._graph.nodes[node_id].state = NodeState.SKIPPED + + def is_node_ready(self, node_id: str) -> bool: + """Check if a node is ready to be executed. + + A node is ready when all its incoming edges from taken branches + have been satisfied. + + Args: + node_id: The ID of the node to check + + Returns: + True if the node is ready for execution + + """ + incoming_edges = self._graph.get_incoming_edges(node_id) + if not incoming_edges: + return True + if any(edge.state == NodeState.UNKNOWN for edge in incoming_edges): + return False + return any(edge.state == NodeState.TAKEN for edge in incoming_edges) + + # ============= Edge State Operations ============= + + def mark_edge_taken(self, edge_id: str) -> None: + """Mark an edge as TAKEN. + + Args: + edge_id: The ID of the edge to mark + + """ + self._graph.edges[edge_id].state = NodeState.TAKEN + + def mark_edge_skipped(self, edge_id: str) -> None: + """Mark an edge as SKIPPED. + + Args: + edge_id: The ID of the edge to mark + + """ + self._graph.edges[edge_id].state = NodeState.SKIPPED + + def analyze_edge_states(self, edges: list[Edge]) -> _EdgeStateAnalysis: + """Analyze the states of edges and return summary flags. + + Args: + edges: List of edges to analyze + + Returns: + Analysis result with state flags + + """ + states = {edge.state for edge in edges} + return _EdgeStateAnalysis( + has_unknown=NodeState.UNKNOWN in states, + has_taken=NodeState.TAKEN in states, + all_skipped=(states == frozenset((NodeState.SKIPPED,)) if states else True), + ) + + def categorize_branch_edges( + self, + node_id: str, + selected_handle: str, + ) -> tuple[Sequence[Edge], Sequence[Edge]]: + """Categorize branch edges into selected and unselected. + + Args: + node_id: The ID of the branch node + selected_handle: The handle of the selected edge + + Returns: + A tuple of (selected_edges, unselected_edges) + + """ + outgoing_edges = self._graph.get_outgoing_edges(node_id) + selected_edges: list[Edge] = [] + unselected_edges: list[Edge] = [] + for edge in outgoing_edges: + if edge.source_handle == selected_handle: + selected_edges.append(edge) + else: + unselected_edges.append(edge) + return selected_edges, unselected_edges + + # ============= Execution Tracking Operations ============= + + def track_unfinished(self, node_id: str) -> None: + """Restore an unfinished node to this frame's execution tracking. + + Args: + node_id: The ID of the unfinished node + + """ + self._unfinished_nodes.add(node_id) + + def finish_execution(self, node_id: str) -> None: + """Mark a node as no longer pending or running. + + Args: + node_id: The ID of the node finishing execution + + """ + self._unfinished_nodes.discard(node_id) + + # ============= Composite Operations ============= + + def is_execution_complete(self) -> bool: + """Check if this frame's execution is complete. + + Tasks are marked executing when they are enqueued, so this frame is + complete when no task in this manager remains pending or running. + + Returns: + True if execution is complete + + """ + return not self._unfinished_nodes + + def defer_ready_tasks(self, tasks: Sequence[ReadyTask]) -> None: + """Move unclaimed tasks into deferred storage.""" + for task in tasks: + self._state.defer_ready_task(task) + + def process_node_success( + self, + node_id: str, + selected_handle: str | None = None, + ) -> tuple[Sequence[str], Sequence[GraphTraversalEvent]]: + """Advance this frame after a node succeeds. + + Branch nodes follow only the selected handle and propagate skipped + paths. Other nodes take every outgoing edge. The returned node IDs are + ready to be enqueued, while the returned events describe every edge + transition in traversal order. + + Args: + node_id: ID of the node that completed successfully. + selected_handle: Selected branch handle for branch nodes. + + Returns: + Ready downstream node IDs and their edge traversal events. + + """ + node = self._graph.nodes[node_id] + if node.execution_type == NodeExecutionType.BRANCH: + return self.handle_branch_completion(node_id, selected_handle) + return self._process_taken_edges(self._graph.get_outgoing_edges(node_id)) + + def _process_taken_edges( + self, + edges: Sequence[Edge], + ) -> tuple[list[str], list[GraphEdgeTakenEvent]]: + """Take each edge and collect downstream nodes that become ready. + + Args: + edges: Outgoing edges selected by the completed node. + + Returns: + Ready downstream node IDs and emitted taken-edge events. + + """ + ready_nodes: list[str] = [] + traversal_events: list[GraphEdgeTakenEvent] = [] + for edge in edges: + nodes, events = self._process_taken_edge(edge) + ready_nodes.extend(nodes) + traversal_events.extend(events) + return ready_nodes, traversal_events + + def _process_taken_edge( + self, + edge: Edge, + ) -> tuple[Sequence[str], Sequence[GraphEdgeTakenEvent]]: + """Take one edge and report whether its target is ready. + + Args: + edge: Edge whose state should transition to ``TAKEN``. + + Returns: + The target node when ready and the corresponding traversal event. + + """ + self.mark_edge_taken(edge.id) + ready_nodes = [edge.head] if self.is_node_ready(edge.head) else [] + return ready_nodes, [self._build_taken_event(edge)] + + def handle_branch_completion( + self, + node_id: str, + selected_handle: str | None, + ) -> tuple[Sequence[str], Sequence[GraphTraversalEvent]]: + """Advance a branch node along its selected path. + + The selected edges are taken and every unselected path is propagated as + skipped. A missing selection is invalid because the scheduler cannot + infer which branch should run. + + Args: + node_id: ID of the completed branch node. + selected_handle: Handle selected by the branch result. + + Returns: + Ready downstream node IDs and all resulting traversal events. + + Raises: + ValueError: If the branch completed without a selected handle. + + """ + if not selected_handle: + msg = f"Branch node {node_id} completed without selecting a branch" + raise ValueError(msg) + + selected_edges, unselected_edges = self.categorize_branch_edges( + node_id, + selected_handle, + ) + skipped_events = self._skip_branch_paths(unselected_edges) + ready_nodes, taken_events = self._process_taken_edges(selected_edges) + return ready_nodes, [*skipped_events, *taken_events] + + def _skip_branch_paths( + self, + unselected_edges: Sequence[Edge], + ) -> list[GraphEdgeSkippedEvent]: + """Skip every path beginning with an unselected branch edge. + + Args: + unselected_edges: Branch edges not selected by the node result. + + Returns: + Skipped-edge events in graph traversal order. + + """ + events: list[GraphEdgeSkippedEvent] = [] + for edge in unselected_edges: + events.extend(self._skip_edge_path(edge)) + return events + + def _skip_edge_path(self, edge: Edge) -> list[GraphEdgeSkippedEvent]: + """Skip one edge and propagate its effect through the target path. + + Args: + edge: Edge whose state should transition to ``SKIPPED``. + + Returns: + This edge's event followed by downstream skipped-edge events. + + """ + self.mark_edge_skipped(edge.id) + return [ + self._build_skipped_event(edge), + *self._propagate_skip_from_edge(edge.id), + ] + + def _propagate_skip_from_edge(self, edge_id: str) -> list[GraphEdgeSkippedEvent]: + """Resolve the target node after one incoming edge is skipped. + + Propagation waits while another incoming edge remains unknown. A taken + incoming edge makes the target executable; otherwise an entirely + skipped input set skips the target and continues through its outputs. + + Args: + edge_id: ID of the edge that was just skipped. + + Returns: + Additional skipped-edge events produced downstream. + + """ + downstream_node_id = self._graph.edges[edge_id].head + incoming_edges = self._graph.get_incoming_edges(downstream_node_id) + edge_states = self.analyze_edge_states(incoming_edges) + + if edge_states["has_unknown"]: + return [] + if edge_states["has_taken"]: + self.enqueue_node(downstream_node_id) + return [] + if edge_states["all_skipped"]: + return self._propagate_skip_to_node(downstream_node_id) + return [] + + def _propagate_skip_to_node(self, node_id: str) -> list[GraphEdgeSkippedEvent]: + """Skip a node and recursively skip each of its outgoing paths. + + Args: + node_id: ID of the node whose inputs are all skipped. + + Returns: + Skipped-edge events produced from the node's outgoing edges. + + """ + self.mark_node_skipped(node_id) + events: list[GraphEdgeSkippedEvent] = [] + for edge in self._graph.get_outgoing_edges(node_id): + events.extend(self._skip_edge_path(edge)) + return events + + @staticmethod + def _build_taken_event(edge: Edge) -> GraphEdgeTakenEvent: + """Build the public traversal event for an edge marked as taken. + + Args: + edge: Taken graph edge to describe. + + Returns: + An event containing the edge identity, endpoints, and source handle. + + """ + return GraphEdgeTakenEvent( + edge_id=edge.id, + source_node_id=edge.tail, + target_node_id=edge.head, + source_handle=edge.source_handle, + ) + + @staticmethod + def _build_skipped_event(edge: Edge) -> GraphEdgeSkippedEvent: + """Build the public traversal event for an edge marked as skipped. + + Args: + edge: Skipped graph edge to describe. + + Returns: + An event containing the edge identity, endpoints, and source handle. + + """ + return GraphEdgeSkippedEvent( + edge_id=edge.id, + source_node_id=edge.tail, + target_node_id=edge.head, + source_handle=edge.source_handle, + ) diff --git a/src/graphon/engine/worker/__init__.py b/src/graphon/engine/worker/__init__.py new file mode 100644 index 00000000..0d4f8d54 --- /dev/null +++ b/src/graphon/engine/worker/__init__.py @@ -0,0 +1,12 @@ +"""Worker threads, dispatch messages, and their fixed-size pool.""" + +from .worker import ContainerAwaitTask, DispatchTask, NodeEventTask, Worker +from .worker_pool import WorkerPool + +__all__ = [ + "ContainerAwaitTask", + "DispatchTask", + "NodeEventTask", + "Worker", + "WorkerPool", +] diff --git a/src/graphon/graph_engine/worker.py b/src/graphon/engine/worker/worker.py similarity index 77% rename from src/graphon/graph_engine/worker.py rename to src/graphon/engine/worker/worker.py index 7f52aead..cd4b474b 100644 --- a/src/graphon/graph_engine/worker.py +++ b/src/graphon/engine/worker/worker.py @@ -1,50 +1,57 @@ -"""Worker - Thread implementation for queue-based node execution +"""Worker thread for queue-based node execution. -Workers pull node IDs from the ready_queue, execute nodes, and push events -to the event_queue for the dispatcher to process. +Workers pull tasks from the ready queue, execute nodes, and push dispatch tasks +to the dispatch queue for the dispatcher to process. """ import logging import queue import threading -import time from collections.abc import Iterator, Sequence from contextlib import AbstractContextManager, nullcontext +from dataclasses import dataclass from datetime import UTC, datetime from typing import final, override from uuid import uuid4 -from graphon.enums import WorkflowNodeExecutionStatus -from graphon.graph_engine.entities.tasks import ( - ContainerAwaitTask, - DispatchTask, - TaskEvent, -) -from graphon.graph_engine.frames import FrameRegistry -from graphon.graph_engine.layers.base import GraphEngineLayer -from graphon.graph_engine.ready_queue import ( - ROOT_FRAME_ID, +from graphon.engine.frame import FrameRegistry +from graphon.engine.layer import Layer +from graphon.engine.ready_queue import ( ReadyQueue, ReadyTask, StartTask, ) -from graphon.graph_events.base import GraphNodeEventBase -from graphon.graph_events.node import ( +from graphon.engine_events.base import NodeEvent +from graphon.engine_events.node import ( NodeRunFailedEvent, NodeRunStartedEvent, is_node_result_event, ) +from graphon.enums import WorkflowNodeExecutionStatus from graphon.node_events.base import NodeRunResult from graphon.nodes.base.node import Node from graphon.nodes.container_effects import ( ContainerAwaitRequest, ) from graphon.runtime.container_state import create_container_run_state +from graphon.runtime.execution import ROOT_FRAME_ID logger = logging.getLogger(__name__) -WORKER_IDLE_THRESHOLD_SECONDS = 0.2 -NodeEventStream = Iterator[GraphNodeEventBase | ContainerAwaitRequest] + +@dataclass(frozen=True, slots=True) +class NodeEventTask: + frame_id: str + event: NodeEvent + + +@dataclass(frozen=True, slots=True) +class ContainerAwaitTask: + invocation_id: str + request: ContainerAwaitRequest + + +type DispatchTask = NodeEventTask | ContainerAwaitTask @final @@ -52,16 +59,16 @@ class Worker(threading.Thread): """Worker thread that executes nodes from the ready queue. Workers continuously pull node IDs from the ready_queue, execute the - corresponding nodes, and push the resulting events to the event_queue + corresponding nodes, and push the resulting tasks to the dispatch queue for the dispatcher to process. """ def __init__( self, ready_queue: ReadyQueue, - event_queue: queue.Queue[DispatchTask], + dispatch_queue: queue.Queue[DispatchTask], frame_registry: FrameRegistry, - layers: Sequence[GraphEngineLayer], + layers: Sequence[Layer], task_claim_lock: threading.Lock, task_claiming: threading.Event, worker_id: int = 0, @@ -71,16 +78,16 @@ def __init__( Args: ready_queue: Ready queue containing node IDs ready for execution - event_queue: Queue for pushing task-scoped execution events + dispatch_queue: Queue for pushing task-scoped execution results. frame_registry: Registry containing frame-local graphs to execute layers: Graph engine layers for node execution hooks worker_id: Unique identifier for this worker execution_context: Optional execution context for context preservation """ - super().__init__(name=f"GraphWorker-{worker_id}", daemon=True) + super().__init__(name=f"EngineWorker-{worker_id}", daemon=True) self._ready_queue = ready_queue - self._event_queue = event_queue + self._dispatch_queue = dispatch_queue self._frame_registry = frame_registry self._execution_context = ( execution_context if execution_context is not None else nullcontext() @@ -89,7 +96,6 @@ def __init__( self._layers = layers self._task_claim_lock = task_claim_lock self._task_claiming = task_claiming - self._last_task_time = time.time() self._current_node_started_at: datetime | None = None self._current_node: Node | None = None self._current_frame_id = ROOT_FRAME_ID @@ -99,19 +105,6 @@ def stop(self) -> None: """Signal the worker to stop processing.""" self._stop_event.set() - @property - def is_idle(self) -> bool: - """Check if the worker is currently idle.""" - return ( - not self._has_current_task.is_set() - and (time.time() - self._last_task_time) > WORKER_IDLE_THRESHOLD_SECONDS - ) - - @property - def idle_duration(self) -> float: - """Get the duration in seconds since the worker last processed a task.""" - return time.time() - self._last_task_time - @property def has_current_task(self) -> bool: """Return True while the worker owns a queue task.""" @@ -122,7 +115,7 @@ def run(self) -> None: """Main worker loop. Continuously pulls node IDs from ready_queue, executes them, - and pushes events to event_queue until stopped. + and pushes results to the dispatch queue until stopped. """ while not self._stop_event.is_set(): with self._task_claim_lock: @@ -133,7 +126,6 @@ def run(self) -> None: except queue.Empty: task_claimed = False else: - self._last_task_time = time.time() self._has_current_task.set() task_claimed = True if not task_claimed: @@ -149,8 +141,8 @@ def run(self) -> None: "Worker failed while executing node %s", node.id, ) - self._event_queue.put( - TaskEvent( + self._dispatch_queue.put( + NodeEventTask( frame_id=self._current_frame_id, event=self._build_fallback_failure_event( node, @@ -169,16 +161,14 @@ def run(self) -> None: def _execute_task(self, task: ReadyTask) -> None: if isinstance(task, StartTask): self._current_frame_id = task.frame_id - node = self._frame_registry.get(task.frame_id).graph.nodes[task.node_id] + node = self._frame_registry[task.frame_id].graph.nodes[task.node_id] self._current_node = node self._execute_node(frame_id=task.frame_id, node=node) return - root_runtime_state = self._frame_registry.get(ROOT_FRAME_ID).graph_runtime_state + root_runtime_state = self._frame_registry[ROOT_FRAME_ID].state run_state = root_runtime_state.get_container_run(task.invocation_id) self._current_frame_id = run_state.frame_id - node = self._frame_registry.get(run_state.frame_id).graph.nodes[ - run_state.node_id - ] + node = self._frame_registry[run_state.frame_id].graph.nodes[run_state.node_id] self._bind_execution_id(frame_id=run_state.frame_id, node=node) self._current_node = node self._current_node_started_at = run_state.started_at @@ -214,12 +204,10 @@ def _execute_node(self, *, frame_id: str, node: Node) -> None: ) def _bind_execution_id(self, *, frame_id: str, node: Node) -> None: - frame = self._frame_registry.get(frame_id) - node_execution = ( - frame.graph_runtime_state.graph_execution.get_or_create_node_execution( - frame_id=frame_id, - node_id=node.id, - ) + frame = self._frame_registry[frame_id] + node_execution = frame.state.graph_execution.get_or_create_node_execution( + frame_id=frame_id, + node_id=node.id, ) node.bind_execution_id(node_execution.execution_id) @@ -228,10 +216,10 @@ def _run_node_events( *, invocation_id: str | None, node: Node, - node_events: NodeEventStream, + node_events: Iterator[NodeEvent | ContainerAwaitRequest], ) -> bool: error: Exception | None = None - result_event: GraphNodeEventBase | None = None + result_event: NodeEvent | None = None suspended = False with self._execution_context: if invocation_id is None: @@ -256,18 +244,16 @@ def _consume_node_events( *, invocation_id: str | None, node: Node, - node_events: NodeEventStream, - ) -> tuple[GraphNodeEventBase | None, bool]: - result_event: GraphNodeEventBase | None = None + node_events: Iterator[NodeEvent | ContainerAwaitRequest], + ) -> tuple[NodeEvent | None, bool]: + result_event: NodeEvent | None = None for event in node_events: if isinstance(event, ContainerAwaitRequest): started_at = self._current_node_started_at if started_at is None: msg = "container await request emitted before node start" raise RuntimeError(msg) - root_runtime_state = self._frame_registry.get( - ROOT_FRAME_ID, - ).graph_runtime_state + root_runtime_state = self._frame_registry[ROOT_FRAME_ID].state if invocation_id is None: invocation_id = str(uuid4()) root_runtime_state.put_container_run( @@ -279,7 +265,7 @@ def _consume_node_events( request=event, ) ) - self._event_queue.put( + self._dispatch_queue.put( ContainerAwaitTask( invocation_id=invocation_id, request=event, @@ -288,8 +274,8 @@ def _consume_node_events( return None, True if isinstance(event, NodeRunStartedEvent) and event.id == node.execution_id: self._current_node_started_at = event.start_at - self._event_queue.put( - TaskEvent(frame_id=self._current_frame_id, event=event) + self._dispatch_queue.put( + NodeEventTask(frame_id=self._current_frame_id, event=event) ) if is_node_result_event(event): result_event = event @@ -312,7 +298,7 @@ def _invoke_node_run_end_hooks( self, node: Node, error: Exception | None, - result_event: GraphNodeEventBase | None = None, + result_event: NodeEvent | None = None, ) -> None: """Invoke on_node_run_end hooks for all layers.""" for layer in self._layers: @@ -340,7 +326,6 @@ def _build_fallback_failure_event( id=node.execution_id, node_id=node.id, node_type=node.node_type, - in_iteration_id=None, error=error_message, start_at=started_at or failure_time, finished_at=failure_time, diff --git a/src/graphon/engine/worker/worker_pool.py b/src/graphon/engine/worker/worker_pool.py new file mode 100644 index 00000000..3673c861 --- /dev/null +++ b/src/graphon/engine/worker/worker_pool.py @@ -0,0 +1,123 @@ +"""Fixed-size worker pool.""" + +import logging +import queue +import threading +from contextlib import AbstractContextManager +from typing import final + +from graphon.engine.frame import FrameRegistry +from graphon.engine.ready_queue import ReadyQueue, ReadyTask + +from ..layer import Layer +from .worker import DispatchTask, Worker + +logger = logging.getLogger(__name__) + + +@final +class WorkerPool: + """Manage the fixed number of workers configured for an engine.""" + + def __init__( + self, + ready_queue: ReadyQueue, + dispatch_queue: queue.Queue[DispatchTask], + frame_registry: FrameRegistry, + layers: list[Layer], + workers: int, + execution_context: AbstractContextManager[object] | None = None, + ) -> None: + """Initialize the fixed-size worker pool. + + Args: + ready_queue: Ready queue protocol for nodes ready for execution + dispatch_queue: Queue for worker dispatch tasks. + frame_registry: Registry containing frame-local graphs to execute + layers: Graph engine layers for node execution hooks + workers: Fixed number of worker threads to create + execution_context: Optional execution context for context preservation + + Raises: + ValueError: If ``workers`` is not a positive integer. + + """ + if not isinstance(workers, int) or isinstance(workers, bool) or workers < 1: + msg = "workers must be a positive integer" + raise ValueError(msg) + + self._ready_queue = ready_queue + self._dispatch_queue = dispatch_queue + self._frame_registry = frame_registry + self._execution_context = execution_context + self._layers = layers + self._worker_count = workers + + # Worker management + self._workers: list[Worker] = [] + self._lock = threading.Lock() + self._task_claim_lock = threading.Lock() + self._task_claiming = threading.Event() + + def start(self) -> None: + """Start the worker pool.""" + with self._lock: + if self._workers: + return + + self._task_claiming.set() + logger.debug("Starting worker pool: %d workers", self._worker_count) + for worker_id in range(self._worker_count): + self._create_worker(worker_id) + + def stop(self) -> None: + """Stop all workers in the pool.""" + with self._lock: + with self._task_claim_lock: + self._task_claiming.clear() + worker_count = len(self._workers) + + if worker_count > 0: + logger.debug("Stopping worker pool: %d workers", worker_count) + + # Stop all workers + for worker in self._workers: + worker.stop() + + # Wait for workers to finish + for worker in self._workers: + if worker.is_alive(): + worker.join(timeout=2.0) + + self._workers.clear() + + def drain(self) -> list[ReadyTask]: + """Atomically stop task claims and remove unclaimed ready work.""" + with self._lock: + with self._task_claim_lock: + self._task_claiming.clear() + tasks = self._ready_queue.drain() + for worker in self._workers: + if not worker.has_current_task: + worker.stop() + return tasks + + def has_current_tasks(self) -> bool: + with self._lock: + return any(worker.has_current_task for worker in self._workers) + + def _create_worker(self, worker_id: int) -> None: + """Create and start a new worker.""" + worker = Worker( + ready_queue=self._ready_queue, + dispatch_queue=self._dispatch_queue, + frame_registry=self._frame_registry, + layers=self._layers, + worker_id=worker_id, + execution_context=self._execution_context, + task_claim_lock=self._task_claim_lock, + task_claiming=self._task_claiming, + ) + + worker.start() + self._workers.append(worker) diff --git a/src/graphon/graph_events/__init__.py b/src/graphon/engine_events/__init__.py similarity index 93% rename from src/graphon/graph_events/__init__.py rename to src/graphon/engine_events/__init__.py index 349f355c..17283f3c 100644 --- a/src/graphon/graph_events/__init__.py +++ b/src/graphon/engine_events/__init__.py @@ -2,11 +2,7 @@ from .agent import NodeRunAgentLogEvent # Base events -from .base import ( - BaseGraphEvent, - GraphEngineEvent, - GraphNodeEventBase, -) +from .base import EngineEvent, NodeEvent # Graph events from .graph import ( @@ -57,17 +53,16 @@ ) __all__ = [ - "BaseGraphEvent", + "EngineEvent", "GraphEdgeSkippedEvent", "GraphEdgeTakenEvent", - "GraphEngineEvent", - "GraphNodeEventBase", "GraphRunAbortedEvent", "GraphRunFailedEvent", "GraphRunPartialSucceededEvent", "GraphRunPausedEvent", "GraphRunStartedEvent", "GraphRunSucceededEvent", + "NodeEvent", "NodeRunAgentLogEvent", "NodeRunExceptionEvent", "NodeRunFailedEvent", diff --git a/src/graphon/graph_events/agent.py b/src/graphon/engine_events/agent.py similarity index 85% rename from src/graphon/graph_events/agent.py rename to src/graphon/engine_events/agent.py index ca355286..9e3dce50 100644 --- a/src/graphon/graph_events/agent.py +++ b/src/graphon/engine_events/agent.py @@ -2,10 +2,10 @@ from pydantic import Field -from .base import GraphAgentNodeEventBase +from .base import NodeEvent -class NodeRunAgentLogEvent(GraphAgentNodeEventBase): +class NodeRunAgentLogEvent(NodeEvent): message_id: str = Field(..., description="message id") label: str = Field(..., description="label") node_execution_id: str = Field(..., description="node execution id") diff --git a/src/graphon/graph_events/base.py b/src/graphon/engine_events/base.py similarity index 52% rename from src/graphon/graph_events/base.py rename to src/graphon/engine_events/base.py index c6940716..a8fdb0f7 100644 --- a/src/graphon/graph_events/base.py +++ b/src/graphon/engine_events/base.py @@ -4,28 +4,19 @@ from graphon.node_events.base import NodeRunResult -class GraphEngineEvent(BaseModel): - pass +class EngineEvent(BaseModel): + """Base model for events emitted by the engine.""" -class BaseGraphEvent(GraphEngineEvent): - pass +class NodeEvent(EngineEvent): + """Engine event associated with one node execution.""" - -class GraphNodeEventBase(GraphEngineEvent): id: str = Field(..., description="node execution id") node_id: str node_type: NodeType - - in_iteration_id: str | None = None - """iteration id if node is in iteration""" - in_loop_id: str | None = None - """loop id if node is in loop""" + container_id: str = "" + """ID of the container that directly owns the event's execution frame.""" # The version of the node, or "1" if not specified. node_version: str = "1" node_run_result: NodeRunResult = Field(default_factory=NodeRunResult) - - -class GraphAgentNodeEventBase(GraphNodeEventBase): - pass diff --git a/src/graphon/graph_events/graph.py b/src/graphon/engine_events/graph.py similarity index 83% rename from src/graphon/graph_events/graph.py rename to src/graphon/engine_events/graph.py index 5eddc4f0..51532445 100644 --- a/src/graphon/graph_events/graph.py +++ b/src/graphon/engine_events/graph.py @@ -1,11 +1,11 @@ from pydantic import Field +from graphon.engine_events.base import EngineEvent from graphon.entities.pause_reason import PauseReason from graphon.entities.workflow_start_reason import WorkflowStartReason -from graphon.graph_events.base import BaseGraphEvent -class GraphRunStartedEvent(BaseGraphEvent): +class GraphRunStartedEvent(EngineEvent): # Reason is emitted for workflow start events and is always set. reason: WorkflowStartReason = Field( default=WorkflowStartReason.INITIAL, @@ -13,7 +13,7 @@ class GraphRunStartedEvent(BaseGraphEvent): ) -class GraphRunSucceededEvent(BaseGraphEvent): +class GraphRunSucceededEvent(EngineEvent): """Event emitted when a run completes successfully with final outputs.""" outputs: dict[str, object] = Field( @@ -22,12 +22,12 @@ class GraphRunSucceededEvent(BaseGraphEvent): ) -class GraphRunFailedEvent(BaseGraphEvent): +class GraphRunFailedEvent(EngineEvent): error: str = Field(..., description="failed reason") exceptions_count: int = Field(description="exception count", default=0) -class GraphRunPartialSucceededEvent(BaseGraphEvent): +class GraphRunPartialSucceededEvent(EngineEvent): """Event emitted when a run finishes with partial success and failures.""" exceptions_count: int = Field(..., description="exception count") @@ -37,7 +37,7 @@ class GraphRunPartialSucceededEvent(BaseGraphEvent): ) -class GraphRunAbortedEvent(BaseGraphEvent): +class GraphRunAbortedEvent(EngineEvent): """Event emitted when a graph run is aborted by user command.""" reason: str | None = Field(default=None, description="reason for abort") @@ -47,7 +47,7 @@ class GraphRunAbortedEvent(BaseGraphEvent): ) -class GraphRunPausedEvent(BaseGraphEvent): +class GraphRunPausedEvent(EngineEvent): """Event emitted when a graph run is paused by user command.""" reasons: list[PauseReason] = Field( diff --git a/src/graphon/graph_events/iteration.py b/src/graphon/engine_events/iteration.py similarity index 80% rename from src/graphon/graph_events/iteration.py rename to src/graphon/engine_events/iteration.py index 2f91b5a8..eecce7c6 100644 --- a/src/graphon/graph_events/iteration.py +++ b/src/graphon/engine_events/iteration.py @@ -3,10 +3,10 @@ from pydantic import Field -from .base import GraphNodeEventBase +from .base import NodeEvent -class NodeRunIterationStartedEvent(GraphNodeEventBase): +class NodeRunIterationStartedEvent(NodeEvent): node_title: str start_at: datetime = Field(..., description="start at") inputs: Mapping[str, object] = Field(default_factory=dict) @@ -14,13 +14,13 @@ class NodeRunIterationStartedEvent(GraphNodeEventBase): predecessor_node_id: str | None = None -class NodeRunIterationNextEvent(GraphNodeEventBase): +class NodeRunIterationNextEvent(NodeEvent): node_title: str index: int = Field(..., description="index") pre_iteration_output: object = None -class NodeRunIterationSucceededEvent(GraphNodeEventBase): +class NodeRunIterationSucceededEvent(NodeEvent): node_title: str start_at: datetime = Field(..., description="start at") inputs: Mapping[str, object] = Field(default_factory=dict) @@ -29,7 +29,7 @@ class NodeRunIterationSucceededEvent(GraphNodeEventBase): steps: int = 0 -class NodeRunIterationFailedEvent(GraphNodeEventBase): +class NodeRunIterationFailedEvent(NodeEvent): node_title: str start_at: datetime = Field(..., description="start at") inputs: Mapping[str, object] = Field(default_factory=dict) diff --git a/src/graphon/graph_events/loop.py b/src/graphon/engine_events/loop.py similarity index 82% rename from src/graphon/graph_events/loop.py rename to src/graphon/engine_events/loop.py index ef9a1510..aa3fad7b 100644 --- a/src/graphon/graph_events/loop.py +++ b/src/graphon/engine_events/loop.py @@ -3,10 +3,10 @@ from pydantic import Field -from .base import GraphNodeEventBase +from .base import NodeEvent -class NodeRunLoopStartedEvent(GraphNodeEventBase): +class NodeRunLoopStartedEvent(NodeEvent): node_title: str start_at: datetime = Field(..., description="start at") inputs: Mapping[str, object] = Field(default_factory=dict) @@ -14,13 +14,13 @@ class NodeRunLoopStartedEvent(GraphNodeEventBase): predecessor_node_id: str | None = None -class NodeRunLoopNextEvent(GraphNodeEventBase): +class NodeRunLoopNextEvent(NodeEvent): node_title: str index: int = Field(..., description="index") pre_loop_output: object = None -class NodeRunLoopSucceededEvent(GraphNodeEventBase): +class NodeRunLoopSucceededEvent(NodeEvent): node_title: str start_at: datetime = Field(..., description="start at") inputs: Mapping[str, object] = Field(default_factory=dict) @@ -29,7 +29,7 @@ class NodeRunLoopSucceededEvent(GraphNodeEventBase): steps: int = 0 -class NodeRunLoopFailedEvent(GraphNodeEventBase): +class NodeRunLoopFailedEvent(NodeEvent): node_title: str start_at: datetime = Field(..., description="start at") inputs: Mapping[str, object] = Field(default_factory=dict) diff --git a/src/graphon/graph_events/node.py b/src/graphon/engine_events/node.py similarity index 86% rename from src/graphon/graph_events/node.py rename to src/graphon/engine_events/node.py index 37d33933..1ff27c5e 100644 --- a/src/graphon/graph_events/node.py +++ b/src/graphon/engine_events/node.py @@ -7,10 +7,10 @@ from graphon.variables.segments import Segment from graphon.variables.variables import Variable -from .base import GraphNodeEventBase +from .base import NodeEvent -class NodeRunStartedEvent(GraphNodeEventBase): +class NodeRunStartedEvent(NodeEvent): node_title: str predecessor_node_id: str | None = None start_at: datetime = Field(..., description="node start time") @@ -20,7 +20,7 @@ class NodeRunStartedEvent(GraphNodeEventBase): provider_id: str = "" -class NodeRunStreamChunkEvent(GraphNodeEventBase): +class NodeRunStreamChunkEvent(NodeEvent): # Spec-compliant fields selector: Sequence[str] = Field( ..., @@ -35,7 +35,7 @@ class NodeRunStreamChunkEvent(GraphNodeEventBase): ) -class NodeRunReasoningChunkEvent(GraphNodeEventBase): +class NodeRunReasoningChunkEvent(NodeEvent): """Graph-level lift of :class:`StreamReasoningEvent`. The selector identifies the source and lets response filters enforce @@ -59,7 +59,7 @@ class NodeRunReasoningChunkEvent(GraphNodeEventBase): ) -class NodeRunModelPollingProgressEvent(GraphNodeEventBase): +class NodeRunModelPollingProgressEvent(NodeEvent): attempt: int = Field(..., ge=0, description="polling check attempt count") last_checked_at: datetime = Field(..., description="last polling check time") next_check_at: datetime | None = Field( @@ -68,7 +68,7 @@ class NodeRunModelPollingProgressEvent(GraphNodeEventBase): ) -class NodeRunRetrieverResourceEvent(GraphNodeEventBase): +class NodeRunRetrieverResourceEvent(NodeEvent): retriever_resources: Sequence[Mapping[str, object]] = Field( ..., description="retriever resources", @@ -76,12 +76,12 @@ class NodeRunRetrieverResourceEvent(GraphNodeEventBase): context: str = Field(..., description="context") -class NodeRunSucceededEvent(GraphNodeEventBase): +class NodeRunSucceededEvent(NodeEvent): start_at: datetime = Field(..., description="node start time") finished_at: datetime | None = Field(default=None, description="node finish time") -class NodeRunVariableUpdatedEvent(GraphNodeEventBase): +class NodeRunVariableUpdatedEvent(NodeEvent): """Request that the engine apply a variable update before downstream observers continue. """ @@ -89,13 +89,13 @@ class NodeRunVariableUpdatedEvent(GraphNodeEventBase): variable: Variable = Field(..., description="Updated variable payload to apply.") -class NodeRunFailedEvent(GraphNodeEventBase): +class NodeRunFailedEvent(NodeEvent): error: str = Field(..., description="error") start_at: datetime = Field(..., description="node start time") finished_at: datetime | None = Field(default=None, description="node finish time") -class NodeRunExceptionEvent(GraphNodeEventBase): +class NodeRunExceptionEvent(NodeEvent): error: str = Field(..., description="error") start_at: datetime = Field(..., description="node start time") finished_at: datetime | None = Field(default=None, description="node finish time") @@ -109,7 +109,7 @@ class NodeRunRetryEvent(NodeRunStartedEvent): ) -class NodeRunHumanInputFormFilledEvent(GraphNodeEventBase): +class NodeRunHumanInputFormFilledEvent(NodeEvent): """Emitted when a HumanInput form is submitted and before the node finishes.""" node_title: str = Field(..., description="HumanInput node title") @@ -131,18 +131,18 @@ class NodeRunHumanInputFormFilledEvent(GraphNodeEventBase): ) -class NodeRunHumanInputFormTimeoutEvent(GraphNodeEventBase): +class NodeRunHumanInputFormTimeoutEvent(NodeEvent): """Emitted when a HumanInput form times out.""" node_title: str = Field(..., description="HumanInput node title") expiration_time: datetime = Field(..., description="Form expiration time") -class NodeRunPauseRequestedEvent(GraphNodeEventBase): +class NodeRunPauseRequestedEvent(NodeEvent): reason: PauseReason = Field(..., description="pause reason") -def is_node_result_event(event: GraphNodeEventBase) -> bool: +def is_node_result_event(event: NodeEvent) -> bool: """Check if an event is a final result event from node execution. A result event indicates the completion of a node execution and contains diff --git a/src/graphon/graph_events/traversal.py b/src/graphon/engine_events/traversal.py similarity index 81% rename from src/graphon/graph_events/traversal.py rename to src/graphon/engine_events/traversal.py index 420ab462..fe26b8dc 100644 --- a/src/graphon/graph_events/traversal.py +++ b/src/graphon/engine_events/traversal.py @@ -1,13 +1,14 @@ from pydantic import Field -from graphon.graph_events.base import BaseGraphEvent +from graphon.engine_events.base import EngineEvent -class _GraphEdgeTraversalEvent(BaseGraphEvent): +class _GraphEdgeTraversalEvent(EngineEvent): edge_id: str = Field(..., description="edge id") source_node_id: str = Field(..., description="source node id") target_node_id: str = Field(..., description="target node id") source_handle: str | None = Field(default=None, description="source handle") + container_id: str = "" class GraphEdgeTakenEvent(_GraphEdgeTraversalEvent): diff --git a/src/graphon/entities/__init__.py b/src/graphon/entities/__init__.py index ef7789c4..52e28524 100644 --- a/src/graphon/entities/__init__.py +++ b/src/graphon/entities/__init__.py @@ -1,10 +1,10 @@ -from .graph_init_params import GraphInitParams +from .graph_init_params import InitParams from .workflow_execution import WorkflowExecution from .workflow_node_execution import WorkflowNodeExecution from .workflow_start_reason import WorkflowStartReason __all__ = [ - "GraphInitParams", + "InitParams", "WorkflowExecution", "WorkflowNodeExecution", "WorkflowStartReason", diff --git a/src/graphon/entities/graph_init_params.py b/src/graphon/entities/graph_init_params.py index 3e95fa1d..01da02d2 100644 --- a/src/graphon/entities/graph_init_params.py +++ b/src/graphon/entities/graph_init_params.py @@ -6,8 +6,8 @@ DIFY_RUN_CONTEXT_KEY = "_dify" -class GraphInitParams(BaseModel): - """GraphInitParams encapsulates the configurations and contextual information +class InitParams(BaseModel): + """InitParams encapsulates the configurations and contextual information that remain constant throughout a single execution of the graph engine. A single execution is defined as follows: as long as the execution has not reached diff --git a/src/graphon/entities/workflow_execution.py b/src/graphon/entities/workflow_execution.py index fa6df41b..b83e5ba6 100644 --- a/src/graphon/entities/workflow_execution.py +++ b/src/graphon/entities/workflow_execution.py @@ -14,6 +14,8 @@ from graphon.enums import WorkflowExecutionStatus, WorkflowType +# TODO: Remove after downstream migration. # ruff:ignore[line-contains-todo, missing-todo-author, missing-todo-link] +# Consumers must read graph runtime events/state instead of WorkflowExecution. class WorkflowExecution(BaseModel): """Domain model for a workflow execution within the graph runtime.""" diff --git a/src/graphon/filters/__init__.py b/src/graphon/filters/__init__.py deleted file mode 100644 index e06a5e5e..00000000 --- a/src/graphon/filters/__init__.py +++ /dev/null @@ -1,15 +0,0 @@ -from graphon.graph_engine.filters import ( - GraphEventFilter, - GraphEventFilterContext, - ResponseStreamFilter, - ResumableGraphEventFilter, - filter_graph_events, -) - -__all__ = [ - "GraphEventFilter", - "GraphEventFilterContext", - "ResponseStreamFilter", - "ResumableGraphEventFilter", - "filter_graph_events", -] diff --git a/src/graphon/graph/graph.py b/src/graphon/graph/graph.py index 2f311824..f8a2ebbb 100644 --- a/src/graphon/graph/graph.py +++ b/src/graphon/graph/graph.py @@ -202,6 +202,155 @@ def _filter_canvas_only_nodes( filtered_node_configs.append(dict(node_config)) return filtered_node_configs + @staticmethod + def _node_container_id(node_config: Mapping[str, Any]) -> str: + """Resolve the ID of the container that directly owns a node. + + ``data.container_id`` is authoritative. An empty or absent owner places + the node in the root graph. During migration, one non-empty legacy + ``iteration_id`` or ``loop_id`` is accepted; nested legacy ownership is + ambiguous and must therefore be expressed with ``container_id``. + + :param node_config: raw node configuration from the workflow graph + + Returns: + The direct container node ID, or ``""`` for a root node. + + Raises: + TypeError: If an ownership field is not a string. + ValueError: If legacy fields name different containers. + + """ + data = node_config.get("data") + if not isinstance(data, Mapping): + return "" + + if "container_id" in data: + value = data["container_id"] + if not isinstance(value, str): + msg = "Node data.container_id must be a string" + raise TypeError(msg) + return value + + # Transitional support for existing single-level Loop/Iteration configs. + legacy_ids: set[str] = set() + for field in ("iteration_id", "loop_id"): + value = data.get(field) + if value is None: + continue + if not isinstance(value, str): + msg = f"Node data.{field} must be a string" + raise TypeError(msg) + if not value: + continue + legacy_ids.add(value) + if len(legacy_ids) > 1: + msg = "Nested container nodes must set data.container_id" + raise ValueError(msg) + return next(iter(legacy_ids), "") + + @classmethod + def _scope_graph_config( + cls, + *, + graph_config: Mapping[str, Any], + node_configs: list[dict[str, Any]], + edge_configs: list[dict[str, Any]], + container_id: str, + ) -> tuple[list[dict[str, Any]], list[dict[str, Any]], dict[str, Any]]: + """Split a workflow graph into one frame's executable and visible scope. + + The direct node and edge lists contain only objects owned by + ``container_id`` and are used to construct the frame's executable + :class:`Graph`. The returned graph config also retains recursively nested + containers so the frame can later construct its children, while excluding + every parent and sibling scope. Edges between different direct owners are + invalid because they would bypass the owning container node. + + :param graph_config: complete config visible to the parent frame + :param node_configs: validated raw node dictionaries from that config + :param edge_configs: validated raw edge dictionaries from that config + :param container_id: direct owner to materialize, or ``""`` for root + + Returns: + Direct nodes, direct edges, and the container subtree config. + + Raises: + ValueError: If scopes are orphaned, cyclic, or joined by an edge. + + """ + container_ids = { + node_id: cls._node_container_id(node_config) + for node_config in node_configs + if isinstance((node_id := node_config.get("id")), str) + } + direct_node_ids = { + node_id + for node_id, owning_container_id in container_ids.items() + if owning_container_id == container_id + } + subtree_node_ids = set(direct_node_ids) + while descendants := { + node_id + for node_id, owning_container_id in container_ids.items() + if owning_container_id in subtree_node_ids + and node_id not in subtree_node_ids + }: + subtree_node_ids.update(descendants) + if not container_id and subtree_node_ids != set(container_ids): + orphan_node_ids = sorted(set(container_ids) - subtree_node_ids) + msg = f"Nodes reference unknown or cyclic containers: {orphan_node_ids}" + raise ValueError(msg) + + direct_edge_configs: list[dict[str, Any]] = [] + subtree_edge_configs: list[dict[str, Any]] = [] + for edge_config in edge_configs: + source = edge_config.get("source") + target = edge_config.get("target") + if not isinstance(source, str) or not isinstance(target, str): + continue + + if ( + source in container_ids + and target in container_ids + and container_ids[source] != container_ids[target] + ): + msg = ( + f"Edge '{source}->{target}' crosses container scopes " + f"'{container_ids[source]}' and '{container_ids[target]}'" + ) + raise ValueError(msg) + + has_unknown_endpoint = ( + source not in container_ids or target not in container_ids + ) + if ( + source in direct_node_ids + or target in direct_node_ids + or (not container_id and has_unknown_endpoint) + ): + direct_edge_configs.append(edge_config) + if ( + source in subtree_node_ids + or target in subtree_node_ids + or (not container_id and has_unknown_endpoint) + ): + subtree_edge_configs.append(edge_config) + + scoped_graph_config = dict(graph_config) + scoped_graph_config["nodes"] = [ + node_config + for node_config in node_configs + if node_config.get("id") in subtree_node_ids + ] + scoped_graph_config["edges"] = subtree_edge_configs + direct_node_configs = [ + node_config + for node_config in node_configs + if node_config.get("id") in direct_node_ids + ] + return direct_node_configs, direct_edge_configs, scoped_graph_config + @classmethod def _promote_fail_branch_nodes(cls, nodes: dict[str, Node]) -> None: """Promote nodes configured with FAIL_BRANCH error strategy @@ -309,6 +458,7 @@ def init( graph_config: Mapping[str, Any], node_factory: NodeFactory, root_node_id: str, + container_id: str = "", skip_validation: bool = False, ) -> Graph: """Initialize a graph with an explicit execution entry point. @@ -316,6 +466,7 @@ def init( :param graph_config: graph config containing nodes and edges :param node_factory: factory for creating node instances from config data :param root_node_id: active root node id + :param container_id: direct container scope to materialize; empty for root Returns: Initialized graph instance rooted at `root_node_id`. @@ -325,13 +476,20 @@ def init( """ # Parse configs - edge_configs = graph_config.get("edges", []) - node_configs = graph_config.get("nodes", []) - - edge_configs = _ListObjectDict.validate_python(edge_configs) - node_configs = _ListObjectDict.validate_python(node_configs) - node_configs = cls._filter_canvas_only_nodes(node_configs) - node_configs = _ListNodeConfigDict.validate_python(node_configs) + edge_configs = _ListObjectDict.validate_python(graph_config.get("edges", [])) + raw_node_configs = _ListObjectDict.validate_python( + graph_config.get("nodes", []), + ) + raw_node_configs = cls._filter_canvas_only_nodes(raw_node_configs) + direct_node_configs, direct_edge_configs, scoped_graph_config = ( + cls._scope_graph_config( + graph_config=graph_config, + node_configs=raw_node_configs, + edge_configs=edge_configs, + container_id=container_id, + ) + ) + node_configs = _ListNodeConfigDict.validate_python(direct_node_configs) if not node_configs: msg = "Graph must have at least one node" @@ -345,10 +503,14 @@ def init( raise ValueError(msg) # Build edges - edges, in_edges, out_edges = cls._build_edges(edge_configs) + edges, in_edges, out_edges = cls._build_edges(direct_edge_configs) # Create node instances nodes = cls._create_node_instances(node_configs_map, node_factory) + for node in nodes.values(): + node.graph_config = scoped_graph_config + if isinstance(node, Node): + node.bind_graph_config(scoped_graph_config) # Promote fail-branch nodes to branch execution type at graph level cls._promote_fail_branch_nodes(nodes) @@ -372,7 +534,7 @@ def init( in_edges=in_edges, out_edges=out_edges, root_node=root_node, - graph_config=graph_config, + graph_config=scoped_graph_config, node_factory=node_factory, ) diff --git a/src/graphon/graph/graph_template.py b/src/graphon/graph/graph_template.py index 7af23533..7cd3f7b3 100644 --- a/src/graphon/graph/graph_template.py +++ b/src/graphon/graph/graph_template.py @@ -6,7 +6,7 @@ class GraphTemplate(BaseModel): """Graph Template for container nodes and subgraph expansion - According to GraphEngine V2 spec, GraphTemplate contains: + According to Engine V2 spec, GraphTemplate contains: - nodes: mapping of node definitions - edges: mapping of edge definitions - root_ids: list of root node IDs diff --git a/src/graphon/graph_engine/__init__.py b/src/graphon/graph_engine/__init__.py deleted file mode 100644 index c78ef26f..00000000 --- a/src/graphon/graph_engine/__init__.py +++ /dev/null @@ -1,22 +0,0 @@ -from .config import GraphEngineConfig -from .container_handlers import ContainerHandler, ContainerHandlerFactory -from .filters import ( - GraphEventFilter, - GraphEventFilterContext, - ResponseStreamFilter, - ResumableGraphEventFilter, - filter_graph_events, -) -from .graph_engine import GraphEngine - -__all__ = [ - "ContainerHandler", - "ContainerHandlerFactory", - "GraphEngine", - "GraphEngineConfig", - "GraphEventFilter", - "GraphEventFilterContext", - "ResponseStreamFilter", - "ResumableGraphEventFilter", - "filter_graph_events", -] diff --git a/src/graphon/graph_engine/_engine_utils.py b/src/graphon/graph_engine/_engine_utils.py deleted file mode 100644 index 18afbc1d..00000000 --- a/src/graphon/graph_engine/_engine_utils.py +++ /dev/null @@ -1,20 +0,0 @@ -import time - - -def get_timestamp() -> float: - """Retrieve a timestamp as a float point numer representing the number of seconds - since the Unix epoch. - - This function is primarily used to measure the execution time - of the workflow engine. - Since workflow execution may be paused and resumed on a different machine, - `time.perf_counter` cannot be used as it is inconsistent across machines. - - To address this, the function uses the wall clock as the time source. - However, it assumes that the clocks of all servers are properly synchronized. - - Returns: - Rounded wall-clock seconds since the Unix epoch. - - """ - return round(time.time()) diff --git a/src/graphon/graph_engine/command_channels/__init__.py b/src/graphon/graph_engine/command_channels/__init__.py deleted file mode 100644 index 670e937c..00000000 --- a/src/graphon/graph_engine/command_channels/__init__.py +++ /dev/null @@ -1,7 +0,0 @@ -"""Command channel implementations for GraphEngine.""" - -from .in_memory_channel import InMemoryChannel -from .protocol import CommandChannel -from .redis_channel import RedisChannel - -__all__ = ["CommandChannel", "InMemoryChannel", "RedisChannel"] diff --git a/src/graphon/graph_engine/command_channels/protocol.py b/src/graphon/graph_engine/command_channels/protocol.py deleted file mode 100644 index 350f6c60..00000000 --- a/src/graphon/graph_engine/command_channels/protocol.py +++ /dev/null @@ -1,42 +0,0 @@ -"""CommandChannel protocol for GraphEngine command communication. - -This protocol defines the interface for sending and receiving commands -to/from a GraphEngine instance, supporting both local and distributed scenarios. -""" - -from abc import abstractmethod -from typing import Protocol - -from ..entities.commands import GraphEngineCommand - - -class CommandChannel(Protocol): - """Protocol for bidirectional command communication with GraphEngine. - - Since each GraphEngine instance processes only one workflow execution, - this channel is dedicated to that single execution. - """ - - @abstractmethod - def fetch_commands(self) -> list[GraphEngineCommand]: - """Fetch pending commands for this GraphEngine instance. - - Called by GraphEngine to poll for commands that need to be processed. - - Returns: - List of pending commands (may be empty) - - """ - ... - - @abstractmethod - def send_command(self, command: GraphEngineCommand) -> None: - """Send a command to be processed by this GraphEngine instance. - - Called by external systems to send control commands to the running workflow. - - Args: - command: The command to send - - """ - ... diff --git a/src/graphon/graph_engine/command_processing/__init__.py b/src/graphon/graph_engine/command_processing/__init__.py deleted file mode 100644 index 86bb36fd..00000000 --- a/src/graphon/graph_engine/command_processing/__init__.py +++ /dev/null @@ -1,19 +0,0 @@ -"""Command processing subsystem for graph engine. - -This package handles external commands sent to the engine -during execution. -""" - -from .command_handlers import ( - AbortCommandHandler, - PauseCommandHandler, - UpdateVariablesCommandHandler, -) -from .command_processor import CommandProcessor - -__all__ = [ - "AbortCommandHandler", - "CommandProcessor", - "PauseCommandHandler", - "UpdateVariablesCommandHandler", -] diff --git a/src/graphon/graph_engine/command_processing/command_handlers.py b/src/graphon/graph_engine/command_processing/command_handlers.py deleted file mode 100644 index 745f8e4b..00000000 --- a/src/graphon/graph_engine/command_processing/command_handlers.py +++ /dev/null @@ -1,71 +0,0 @@ -import logging -from typing import final, override - -from graphon.entities.pause_reason import SchedulingPause -from graphon.runtime.graph_runtime_state import GraphExecutionProtocol -from graphon.runtime.variable_pool import VariablePool - -from ..entities.commands import ( - AbortCommand, - PauseCommand, - UpdateVariablesCommand, -) -from .command_processor import CommandHandler - -logger = logging.getLogger(__name__) - - -@final -class AbortCommandHandler(CommandHandler[AbortCommand]): - @override - def handle( - self, - command: AbortCommand, - execution: GraphExecutionProtocol, - ) -> None: - logger.debug("Aborting workflow %s: %s", execution.workflow_id, command.reason) - execution.abort(command.reason or "User requested abort") - - -@final -class PauseCommandHandler(CommandHandler[PauseCommand]): - @override - def handle( - self, - command: PauseCommand, - execution: GraphExecutionProtocol, - ) -> None: - logger.debug("Pausing workflow %s: %s", execution.workflow_id, command.reason) - # Convert string reason to PauseReason if needed - reason = command.reason - pause_reason = SchedulingPause(message=reason) - execution.pause(pause_reason) - - -@final -class UpdateVariablesCommandHandler(CommandHandler[UpdateVariablesCommand]): - def __init__(self, variable_pool: VariablePool) -> None: - self._variable_pool = variable_pool - - @override - def handle( - self, - command: UpdateVariablesCommand, - execution: GraphExecutionProtocol, - ) -> None: - for update in command.updates: - try: - variable = update.value - self._variable_pool.add(variable.selector, variable) - logger.debug( - "Updated variable %s for workflow %s", - variable.selector, - execution.workflow_id, - ) - except ValueError as exc: - logger.warning( - "Skipping invalid variable selector %s for workflow %s: %s", - getattr(update.value, "selector", None), - execution.workflow_id, - exc, - ) diff --git a/src/graphon/graph_engine/command_processing/command_processor.py b/src/graphon/graph_engine/command_processing/command_processor.py deleted file mode 100644 index 1cadb502..00000000 --- a/src/graphon/graph_engine/command_processing/command_processor.py +++ /dev/null @@ -1,110 +0,0 @@ -"""Main command processor for handling external commands.""" - -import logging -from abc import abstractmethod -from collections.abc import Callable -from typing import Protocol, final - -from graphon.runtime.graph_runtime_state import GraphExecutionProtocol - -from ..command_channels import CommandChannel -from ..entities.commands import GraphEngineCommand - -logger = logging.getLogger(__name__) - - -class CommandHandler[CommandT: GraphEngineCommand](Protocol): - """Protocol for command handlers.""" - - @abstractmethod - def handle( - self, - command: CommandT, - execution: GraphExecutionProtocol, - ) -> None: ... - - -@final -class CommandProcessor: - """Processes external commands sent to the engine. - - This polls the command channel and dispatches commands to - appropriate handlers. - """ - - def __init__( - self, - command_channel: CommandChannel, - graph_execution: GraphExecutionProtocol, - ) -> None: - """Initialize the command processor. - - Args: - command_channel: Channel for receiving commands - graph_execution: Graph execution aggregate - - """ - self._command_channel = command_channel - self._graph_execution = graph_execution - self._handlers: dict[ - type[GraphEngineCommand], - Callable[[GraphEngineCommand, GraphExecutionProtocol], None], - ] = {} - - def register_handler[CommandT: GraphEngineCommand]( - self, - command_type: type[CommandT], - handler: CommandHandler[CommandT], - ) -> None: - """Register a handler for a command type. - - Args: - command_type: Type of command to handle - handler: Handler for the command - - """ - - def invoke( - command: GraphEngineCommand, - execution: GraphExecutionProtocol, - ) -> None: - if not isinstance(command, command_type): - msg = ( - f"Registered handler for {command_type.__name__} received " - f"{type(command).__name__}" - ) - raise TypeError(msg) - handler.handle(command, execution) - - self._handlers[command_type] = invoke - - def process_commands(self) -> None: - """Check for and process any pending commands.""" - try: - commands = self._command_channel.fetch_commands() - for command in commands: - self._handle_command(command) - except Exception: - logger.exception("Error processing commands") - - def _handle_command(self, command: GraphEngineCommand) -> None: - """Handle a single command. - - Args: - command: The command to handle - - """ - handler = self._handlers.get(type(command)) - if handler: - try: - handler(command, self._graph_execution) - except Exception: - logger.exception( - "Error handling command %s", - command.__class__.__name__, - ) - else: - logger.warning( - "No handler registered for command: %s", - command.__class__.__name__, - ) diff --git a/src/graphon/graph_engine/config.py b/src/graphon/graph_engine/config.py deleted file mode 100644 index 677f2cfd..00000000 --- a/src/graphon/graph_engine/config.py +++ /dev/null @@ -1,14 +0,0 @@ -"""GraphEngine configuration models.""" - -from pydantic import BaseModel, ConfigDict - - -class GraphEngineConfig(BaseModel): - """Configuration for GraphEngine worker pool scaling.""" - - model_config = ConfigDict(frozen=True) - - min_workers: int = 1 - max_workers: int = 5 - scale_up_threshold: int = 0 - scale_down_idle_time: float = 5.0 diff --git a/src/graphon/graph_engine/container_handlers.py b/src/graphon/graph_engine/container_handlers.py deleted file mode 100644 index 0da9796e..00000000 --- a/src/graphon/graph_engine/container_handlers.py +++ /dev/null @@ -1,57 +0,0 @@ -from __future__ import annotations - -from abc import abstractmethod -from collections.abc import Callable -from typing import Protocol - -from graphon.enums import NodeType -from graphon.graph_events.base import GraphNodeEventBase -from graphon.graph_events.node import NodeRunFailedEvent -from graphon.nodes.container_effects import ContainerAwaitRequest -from graphon.runtime.container_state import ContainerFrameState - -from .frames import ExecutionFrame, FrameRegistry - - -class ContainerHandler(Protocol): - node_type: NodeType - - @abstractmethod - def restore_frame(self, frame_state: ContainerFrameState) -> None: ... - - @abstractmethod - def start_await( - self, - *, - invocation_id: str, - request: ContainerAwaitRequest, - ) -> None: ... - - @abstractmethod - def prepare_frame_event( - self, - *, - frame: ExecutionFrame, - event: GraphNodeEventBase, - ) -> None: ... - - @abstractmethod - def should_collect( - self, - *, - event: GraphNodeEventBase, - ) -> bool: ... - - @abstractmethod - def record_frame_failure( - self, - *, - frame: ExecutionFrame, - event: NodeRunFailedEvent, - ) -> None: ... - - @abstractmethod - def complete_frame(self, frame: ExecutionFrame) -> None: ... - - -type ContainerHandlerFactory = Callable[[FrameRegistry], ContainerHandler] diff --git a/src/graphon/graph_engine/domain/__init__.py b/src/graphon/graph_engine/domain/__init__.py deleted file mode 100644 index 55b9ec3c..00000000 --- a/src/graphon/graph_engine/domain/__init__.py +++ /dev/null @@ -1,13 +0,0 @@ -"""Domain models for graph engine. - -This package contains the core domain entities, value objects, and aggregates -that represent the business concepts of workflow graph execution. -""" - -from .graph_execution import GraphExecution -from .node_execution import NodeExecution - -__all__ = [ - "GraphExecution", - "NodeExecution", -] diff --git a/src/graphon/graph_engine/domain/node_execution.py b/src/graphon/graph_engine/domain/node_execution.py deleted file mode 100644 index 1cfa6746..00000000 --- a/src/graphon/graph_engine/domain/node_execution.py +++ /dev/null @@ -1,19 +0,0 @@ -"""NodeExecution entity representing a node's execution state.""" - -from dataclasses import dataclass - - -@dataclass -class NodeExecution: - """Entity representing the execution state of a single node. - - This is a mutable entity that tracks the runtime state of a node - during graph execution. - """ - - execution_id: str - retry_count: int = 0 - - def increment_retry(self) -> None: - """Increment the retry count for this node.""" - self.retry_count += 1 diff --git a/src/graphon/graph_engine/entities/commands.py b/src/graphon/graph_engine/entities/commands.py deleted file mode 100644 index 2827fc3d..00000000 --- a/src/graphon/graph_engine/entities/commands.py +++ /dev/null @@ -1,70 +0,0 @@ -"""GraphEngine command entities for external control. - -This module defines command types that can be sent to a running GraphEngine -instance to control its execution flow. -""" - -from collections.abc import Sequence -from enum import StrEnum, auto -from typing import Any - -from pydantic import BaseModel, Field - -from graphon.variables.variables import Variable - - -class CommandType(StrEnum): - """Types of commands that can be sent to GraphEngine.""" - - ABORT = auto() - PAUSE = auto() - UPDATE_VARIABLES = auto() - - -class GraphEngineCommand(BaseModel): - """Base class for all GraphEngine commands.""" - - command_type: CommandType = Field(..., description="Type of command") - payload: dict[str, Any] | None = Field( - default=None, - description="Optional command payload", - ) - - -class AbortCommand(GraphEngineCommand): - """Command to abort a running workflow execution.""" - - command_type: CommandType = Field( - default=CommandType.ABORT, - description="Type of command", - ) - reason: str | None = Field(default=None, description="Optional reason for abort") - - -class PauseCommand(GraphEngineCommand): - """Command to pause a running workflow execution.""" - - command_type: CommandType = Field( - default=CommandType.PAUSE, - description="Type of command", - ) - reason: str = Field(default="unknown reason", description="reason for pause") - - -class VariableUpdate(BaseModel): - """Represents a single variable update instruction.""" - - value: Variable = Field(description="New variable value") - - -class UpdateVariablesCommand(GraphEngineCommand): - """Command to update a group of variables in the variable pool.""" - - command_type: CommandType = Field( - default=CommandType.UPDATE_VARIABLES, - description="Type of command", - ) - updates: Sequence[VariableUpdate] = Field( - default_factory=list, - description="Variable updates", - ) diff --git a/src/graphon/graph_engine/entities/tasks.py b/src/graphon/graph_engine/entities/tasks.py deleted file mode 100644 index 1d2f57af..00000000 --- a/src/graphon/graph_engine/entities/tasks.py +++ /dev/null @@ -1,21 +0,0 @@ -from __future__ import annotations - -from dataclasses import dataclass - -from graphon.graph_events.base import GraphNodeEventBase -from graphon.nodes.container_effects import ContainerAwaitRequest - - -@dataclass(frozen=True, slots=True) -class TaskEvent: - frame_id: str - event: GraphNodeEventBase - - -@dataclass(frozen=True, slots=True) -class ContainerAwaitTask: - invocation_id: str - request: ContainerAwaitRequest - - -type DispatchTask = TaskEvent | ContainerAwaitTask diff --git a/src/graphon/graph_engine/event_management/__init__.py b/src/graphon/graph_engine/event_management/__init__.py deleted file mode 100644 index a4010d58..00000000 --- a/src/graphon/graph_engine/event_management/__init__.py +++ /dev/null @@ -1,13 +0,0 @@ -"""Event management subsystem for graph engine. - -This package handles event routing, collection, and emission for -workflow graph execution events. -""" - -from .event_handlers import EventHandler -from .event_manager import EventManager - -__all__ = [ - "EventHandler", - "EventManager", -] diff --git a/src/graphon/graph_engine/filters/__init__.py b/src/graphon/graph_engine/filters/__init__.py deleted file mode 100644 index e9cdec53..00000000 --- a/src/graphon/graph_engine/filters/__init__.py +++ /dev/null @@ -1,15 +0,0 @@ -from graphon.graph_engine.filters.base import ( - GraphEventFilter, - GraphEventFilterContext, - ResumableGraphEventFilter, -) -from graphon.graph_engine.filters.chain import filter_graph_events -from graphon.graph_engine.filters.response_stream import ResponseStreamFilter - -__all__ = [ - "GraphEventFilter", - "GraphEventFilterContext", - "ResponseStreamFilter", - "ResumableGraphEventFilter", - "filter_graph_events", -] diff --git a/src/graphon/graph_engine/filters/base.py b/src/graphon/graph_engine/filters/base.py deleted file mode 100644 index 5e558979..00000000 --- a/src/graphon/graph_engine/filters/base.py +++ /dev/null @@ -1,70 +0,0 @@ -from __future__ import annotations - -from abc import abstractmethod -from collections.abc import Iterable -from dataclasses import dataclass -from typing import TYPE_CHECKING, Protocol - -from graphon.graph.graph import Graph -from graphon.graph_events.base import GraphEngineEvent -from graphon.runtime.graph_runtime_state_protocol import ReadOnlyGraphRuntimeState -from graphon.runtime.read_only_wrappers import ReadOnlyGraphRuntimeStateWrapper - -if TYPE_CHECKING: - from graphon.graph_engine.graph_engine import GraphEngine - - -@dataclass(frozen=True) -class GraphEventFilterContext: - """Run-scoped context available to graph event filters.""" - - graph: Graph - runtime_state: ReadOnlyGraphRuntimeState - - @classmethod - def from_engine(cls, engine: GraphEngine) -> GraphEventFilterContext: - return cls( - graph=engine.graph, - runtime_state=ReadOnlyGraphRuntimeStateWrapper( - engine.graph_runtime_state, - ), - ) - - -class GraphEventFilter(Protocol): - """Event-to-event transform used outside GraphEngine execution.""" - - @property - @abstractmethod - def filter_id(self) -> str: - """Stable identifier for diagnostics and external state storage.""" - raise NotImplementedError - - @abstractmethod - def initialize(self, context: GraphEventFilterContext) -> None: - """Bind run-scoped context before events are processed.""" - raise NotImplementedError - - @abstractmethod - def on_event(self, event: GraphEngineEvent) -> Iterable[GraphEngineEvent]: - """Transform one input event into zero or more output events.""" - raise NotImplementedError - - @abstractmethod - def flush(self) -> Iterable[GraphEngineEvent]: - """Emit buffered events after the upstream source is exhausted.""" - raise NotImplementedError - - -class ResumableGraphEventFilter(GraphEventFilter, Protocol): - """Optional filter protocol for output-layer resume state.""" - - @abstractmethod - def dumps(self) -> str: - """Serialize this filter's private state.""" - raise NotImplementedError - - @abstractmethod - def loads(self, data: str) -> None: - """Restore this filter's private state.""" - raise NotImplementedError diff --git a/src/graphon/graph_engine/frames.py b/src/graphon/graph_engine/frames.py deleted file mode 100644 index 3aa5b4cb..00000000 --- a/src/graphon/graph_engine/frames.py +++ /dev/null @@ -1,146 +0,0 @@ -"""Execution frame registry for frame-scoped graph tasks.""" - -from __future__ import annotations - -from abc import abstractmethod -from dataclasses import dataclass -from typing import Protocol, cast, final - -from graphon.graph.graph import Graph, NodeFactory -from graphon.runtime.container_state import ContainerFrameState -from graphon.runtime.graph_runtime_state import GraphRuntimeState -from graphon.runtime.variable_pool import VariablePool - -from .error_handler import ErrorHandler -from .graph_state_manager import GraphStateManager -from .graph_traversal.edge_processor import EdgeProcessor -from .graph_traversal.skip_propagator import SkipPropagator -from .ready_queue import ROOT_FRAME_ID - - -class RebindableNodeFactory(NodeFactory, Protocol): - @abstractmethod - def with_runtime_state( - self, - graph_runtime_state: GraphRuntimeState, - ) -> RebindableNodeFactory: ... - - -@dataclass(frozen=True, slots=True) -class ExecutionFrame: - frame_id: str - graph: Graph - graph_runtime_state: GraphRuntimeState - state_manager: GraphStateManager - edge_processor: EdgeProcessor - error_handler: ErrorHandler - - -@final -class FrameRegistry: - def __init__(self) -> None: - self._frames: dict[str, ExecutionFrame] = {} - - def register(self, frame: ExecutionFrame) -> None: - self._frames[frame.frame_id] = frame - - def get(self, frame_id: str) -> ExecutionFrame: - return self._frames[frame_id] - - def remove(self, frame_id: str) -> None: - del self._frames[frame_id] - - def materialize_child_frame( - self, - *, - frame_id: str, - root_node_id: str, - graph_runtime_state: GraphRuntimeState, - ) -> ExecutionFrame: - root_graph = self.get(ROOT_FRAME_ID).graph - graph_config = root_graph.graph_config - if graph_config is None: - msg = "Root graph does not carry graph_config for frame materialization." - raise RuntimeError(msg) - node_factory = root_graph.node_factory - if node_factory is None: - msg = "Root graph does not carry node_factory for frame materialization." - raise RuntimeError(msg) - - rebound_factory = cast(RebindableNodeFactory, node_factory).with_runtime_state( - graph_runtime_state, - ) - graph = Graph.init( - graph_config=graph_config, - node_factory=rebound_factory, - root_node_id=root_node_id, - ) - graph_runtime_state.attach_graph(graph) - state_manager = GraphStateManager( - graph, - graph_runtime_state, - frame_id, - ) - skip_propagator = SkipPropagator( - graph=graph, - state_manager=state_manager, - ) - edge_processor = EdgeProcessor( - graph=graph, - state_manager=state_manager, - skip_propagator=skip_propagator, - ) - frame = ExecutionFrame( - frame_id=frame_id, - graph=graph, - graph_runtime_state=graph_runtime_state, - state_manager=state_manager, - edge_processor=edge_processor, - error_handler=ErrorHandler(graph, graph_runtime_state.graph_execution), - ) - self.register(frame) - return frame - - def materialize_child_frame_from_state( - self, - frame_state: ContainerFrameState, - *, - variable_pool: VariablePool, - ) -> ExecutionFrame: - runtime_data = frame_state.runtime_data - root_runtime_state = self.get(ROOT_FRAME_ID).graph_runtime_state - graph_runtime_state = GraphRuntimeState( - variable_pool=variable_pool, - start_at=root_runtime_state.start_at, - llm_usage=runtime_data.llm_usage, - outputs=dict(runtime_data.outputs), - node_run_steps=runtime_data.node_run_steps, - ready_queue=root_runtime_state.ready_queue, - deferred_ready_queue=root_runtime_state.deferred_ready_queue, - graph_execution=root_runtime_state.graph_execution, - ) - frame = self.materialize_child_frame( - frame_id=frame_state.frame_id, - root_node_id=frame_state.root_node_id, - graph_runtime_state=graph_runtime_state, - ) - missing_node_ids = sorted( - set(runtime_data.graph_node_states) - set(frame.graph.nodes), - ) - missing_edge_ids = sorted( - set(runtime_data.graph_edge_states) - set(frame.graph.edges), - ) - if missing_node_ids or missing_edge_ids: - msg = ( - f"Saved frame state for {frame_state.frame_id} does not match " - f"rebuilt graph: missing node ids={missing_node_ids}, " - f"missing edge ids={missing_edge_ids}" - ) - self.remove(frame_state.frame_id) - raise RuntimeError(msg) - - for node_id, state in runtime_data.graph_node_states.items(): - frame.graph.nodes[node_id].state = state - for edge_id, state in runtime_data.graph_edge_states.items(): - frame.graph.edges[edge_id].state = state - return frame diff --git a/src/graphon/graph_engine/graph_state_manager.py b/src/graphon/graph_engine/graph_state_manager.py deleted file mode 100644 index d24a2a1c..00000000 --- a/src/graphon/graph_engine/graph_state_manager.py +++ /dev/null @@ -1,214 +0,0 @@ -"""Graph state manager that combines node, edge, and execution tracking.""" - -import threading -from collections.abc import Sequence -from typing import TypedDict, final - -from graphon.enums import NodeState -from graphon.graph.edge import Edge -from graphon.graph.graph import Graph -from graphon.runtime.graph_runtime_state import GraphRuntimeState - -from .ready_queue import ReadyTask, StartTask - - -class EdgeStateAnalysis(TypedDict): - """Analysis result for edge states.""" - - has_unknown: bool - has_taken: bool - all_skipped: bool - - -@final -class GraphStateManager: - def __init__( - self, - graph: Graph, - graph_runtime_state: GraphRuntimeState, - frame_id: str, - ) -> None: - """Initialize the state manager. - - Args: - graph: The workflow graph - graph_runtime_state: Runtime state owning ready task queues - frame_id: Execution frame managed by this instance - - """ - self._graph = graph - self._graph_runtime_state = graph_runtime_state - self._frame_id = frame_id - self._lock = threading.Lock() - - self._unfinished_nodes: set[str] = set() - - # ============= Node State Operations ============= - - def enqueue_node(self, node_id: str) -> None: - """Mark a node as TAKEN and add its task to the ready queue. - - This combines the state transition and enqueueing operations - that always occur together when preparing a node for execution. - - Args: - node_id: The ID of the node to enqueue - - """ - with self._lock: - self._graph.nodes[node_id].state = NodeState.TAKEN - self._unfinished_nodes.add(node_id) - self._graph_runtime_state.enqueue_ready_task( - StartTask(frame_id=self._frame_id, node_id=node_id), - ) - - def mark_node_skipped(self, node_id: str) -> None: - """Mark a node as SKIPPED. - - Args: - node_id: The ID of the node to skip - - """ - with self._lock: - self._graph.nodes[node_id].state = NodeState.SKIPPED - - def is_node_ready(self, node_id: str) -> bool: - """Check if a node is ready to be executed. - - A node is ready when all its incoming edges from taken branches - have been satisfied. - - Args: - node_id: The ID of the node to check - - Returns: - True if the node is ready for execution - - """ - with self._lock: - # Get all incoming edges to this node - incoming_edges = self._graph.get_incoming_edges(node_id) - - # If no incoming edges, node is always ready - if not incoming_edges: - return True - - # If any edge is UNKNOWN, node is not ready - if any(edge.state == NodeState.UNKNOWN for edge in incoming_edges): - return False - - # Node is ready if at least one edge is TAKEN - return any(edge.state == NodeState.TAKEN for edge in incoming_edges) - - # ============= Edge State Operations ============= - - def mark_edge_taken(self, edge_id: str) -> None: - """Mark an edge as TAKEN. - - Args: - edge_id: The ID of the edge to mark - - """ - with self._lock: - self._graph.edges[edge_id].state = NodeState.TAKEN - - def mark_edge_skipped(self, edge_id: str) -> None: - """Mark an edge as SKIPPED. - - Args: - edge_id: The ID of the edge to mark - - """ - with self._lock: - self._graph.edges[edge_id].state = NodeState.SKIPPED - - def analyze_edge_states(self, edges: list[Edge]) -> EdgeStateAnalysis: - """Analyze the states of edges and return summary flags. - - Args: - edges: List of edges to analyze - - Returns: - Analysis result with state flags - - """ - with self._lock: - states = {edge.state for edge in edges} - - return EdgeStateAnalysis( - has_unknown=NodeState.UNKNOWN in states, - has_taken=NodeState.TAKEN in states, - all_skipped=( - states == frozenset((NodeState.SKIPPED,)) if states else True - ), - ) - - def categorize_branch_edges( - self, - node_id: str, - selected_handle: str, - ) -> tuple[Sequence[Edge], Sequence[Edge]]: - """Categorize branch edges into selected and unselected. - - Args: - node_id: The ID of the branch node - selected_handle: The handle of the selected edge - - Returns: - A tuple of (selected_edges, unselected_edges) - - """ - with self._lock: - outgoing_edges = self._graph.get_outgoing_edges(node_id) - selected_edges: list[Edge] = [] - unselected_edges: list[Edge] = [] - - for edge in outgoing_edges: - if edge.source_handle == selected_handle: - selected_edges.append(edge) - else: - unselected_edges.append(edge) - - return selected_edges, unselected_edges - - # ============= Execution Tracking Operations ============= - - def track_unfinished(self, node_id: str) -> None: - """Restore an unfinished node to this frame's execution tracking. - - Args: - node_id: The ID of the unfinished node - - """ - with self._lock: - self._unfinished_nodes.add(node_id) - - def finish_execution(self, node_id: str) -> None: - """Mark a node as no longer pending or running. - - Args: - node_id: The ID of the node finishing execution - - """ - with self._lock: - self._unfinished_nodes.discard(node_id) - - # ============= Composite Operations ============= - - def is_execution_complete(self) -> bool: - """Check if this frame's execution is complete. - - Tasks are marked executing when they are enqueued, so this frame is - complete when no task in this manager remains pending or running. - - Returns: - True if execution is complete - - """ - with self._lock: - return not self._unfinished_nodes - - def defer_ready_tasks(self, tasks: Sequence[ReadyTask]) -> None: - """Move unclaimed tasks into deferred storage.""" - for task in tasks: - self._graph_runtime_state.defer_ready_task(task) diff --git a/src/graphon/graph_engine/graph_traversal/__init__.py b/src/graphon/graph_engine/graph_traversal/__init__.py deleted file mode 100644 index b5458e02..00000000 --- a/src/graphon/graph_engine/graph_traversal/__init__.py +++ /dev/null @@ -1,13 +0,0 @@ -"""Graph traversal subsystem for graph engine. - -This package handles graph navigation, edge processing, -and skip propagation logic. -""" - -from .edge_processor import EdgeProcessor -from .skip_propagator import SkipPropagator - -__all__ = [ - "EdgeProcessor", - "SkipPropagator", -] diff --git a/src/graphon/graph_engine/graph_traversal/edge_processor.py b/src/graphon/graph_engine/graph_traversal/edge_processor.py deleted file mode 100644 index da01b81a..00000000 --- a/src/graphon/graph_engine/graph_traversal/edge_processor.py +++ /dev/null @@ -1,174 +0,0 @@ -"""Edge processing logic for graph traversal.""" - -from collections.abc import Sequence -from typing import final - -from graphon.enums import NodeExecutionType -from graphon.graph.edge import Edge -from graphon.graph.graph import Graph -from graphon.graph_events.traversal import GraphEdgeSkippedEvent, GraphEdgeTakenEvent - -from ..graph_state_manager import GraphStateManager -from .skip_propagator import SkipPropagator - -type GraphTraversalEvent = GraphEdgeTakenEvent | GraphEdgeSkippedEvent - - -@final -class EdgeProcessor: - """Processes edges during graph execution. - - This handles marking edges as taken or skipped, emitting traversal events, - triggering downstream node execution, and managing branch node logic. - """ - - def __init__( - self, - graph: Graph, - state_manager: GraphStateManager, - skip_propagator: SkipPropagator, - ) -> None: - """Initialize the edge processor. - - Args: - graph: The workflow graph - state_manager: Unified state manager - skip_propagator: Propagator for skip states - - """ - self._graph = graph - self._state_manager = state_manager - self._skip_propagator = skip_propagator - - def process_node_success( - self, - node_id: str, - selected_handle: str | None = None, - ) -> tuple[Sequence[str], Sequence[GraphTraversalEvent]]: - """Process edges after a node succeeds. - - Args: - node_id: The ID of the succeeded node - selected_handle: For branch nodes, the selected edge handle - - Returns: - Tuple of (list of downstream node IDs that are now ready, - list of traversal events) - - """ - node = self._graph.nodes[node_id] - - if node.execution_type == NodeExecutionType.BRANCH: - return self.handle_branch_completion(node_id, selected_handle) - return self._process_non_branch_node_edges(node_id) - - def _process_non_branch_node_edges( - self, - node_id: str, - ) -> tuple[Sequence[str], Sequence[GraphTraversalEvent]]: - """Process edges for non-branch nodes (mark all as TAKEN). - - Args: - node_id: The ID of the succeeded node - - Returns: - Tuple of (list of downstream nodes ready for execution, - list of traversal events) - - """ - return self._process_taken_edges(self._graph.get_outgoing_edges(node_id)) - - def _process_taken_edges( - self, - edges: Sequence[Edge], - ) -> tuple[list[str], list[GraphEdgeTakenEvent]]: - ready_nodes: list[str] = [] - traversal_events: list[GraphEdgeTakenEvent] = [] - for edge in edges: - nodes, events = self._process_taken_edge(edge) - ready_nodes.extend(nodes) - traversal_events.extend(events) - return ready_nodes, traversal_events - - def _process_taken_edge( - self, - edge: Edge, - ) -> tuple[Sequence[str], Sequence[GraphEdgeTakenEvent]]: - """Mark edge as taken and check downstream node. - - Args: - edge: The edge to process - - Returns: - Tuple of ( - list containing downstream node ID if it's ready, - list of traversal events - ) - - """ - # Mark edge as taken - self._state_manager.mark_edge_taken(edge.id) - - # Check if downstream node is ready - ready_nodes: list[str] = [] - if self._state_manager.is_node_ready(edge.head): - ready_nodes.append(edge.head) - - return ready_nodes, [self._build_taken_event(edge)] - - def handle_branch_completion( - self, - node_id: str, - selected_handle: str | None, - ) -> tuple[Sequence[str], Sequence[GraphTraversalEvent]]: - """Handle completion of a branch node. - - Args: - node_id: The ID of the branch node - selected_handle: The handle of the selected branch - - Returns: - Tuple of (list of downstream nodes ready for execution, - list of traversal events) - - Raises: - ValueError: If no branch was selected - - """ - if not selected_handle: - msg = f"Branch node {node_id} completed without selecting a branch" - raise ValueError(msg) - - selected_edges, unselected_edges = self._state_manager.categorize_branch_edges( - node_id, - selected_handle, - ) - - skipped_events = self._skip_propagator.skip_branch_paths(unselected_edges) - - ready_nodes, taken_events = self._process_taken_edges(selected_edges) - return ready_nodes, [*skipped_events, *taken_events] - - def validate_branch_selection(self, node_id: str, selected_handle: str) -> bool: - """Validate that a branch selection is valid. - - Args: - node_id: The ID of the branch node - selected_handle: The handle to validate - - Returns: - True if the selection is valid - - """ - outgoing_edges = self._graph.get_outgoing_edges(node_id) - valid_handles = {edge.source_handle for edge in outgoing_edges} - return selected_handle in valid_handles - - @staticmethod - def _build_taken_event(edge: Edge) -> GraphEdgeTakenEvent: - return GraphEdgeTakenEvent( - edge_id=edge.id, - source_node_id=edge.tail, - target_node_id=edge.head, - source_handle=edge.source_handle, - ) diff --git a/src/graphon/graph_engine/graph_traversal/skip_propagator.py b/src/graphon/graph_engine/graph_traversal/skip_propagator.py deleted file mode 100644 index 1fa92b7c..00000000 --- a/src/graphon/graph_engine/graph_traversal/skip_propagator.py +++ /dev/null @@ -1,128 +0,0 @@ -"""Skip state propagation through the graph.""" - -from collections.abc import Sequence -from typing import final - -from graphon.graph.edge import Edge -from graphon.graph.graph import Graph -from graphon.graph_events.traversal import GraphEdgeSkippedEvent - -from ..graph_state_manager import GraphStateManager - - -@final -class SkipPropagator: - """Propagates skip states through the graph. - - When a node is skipped, this ensures all downstream nodes - that depend solely on it are also skipped. - """ - - def __init__( - self, - graph: Graph, - state_manager: GraphStateManager, - ) -> None: - """Initialize the skip propagator. - - Args: - graph: The workflow graph - state_manager: Unified state manager - - """ - self._graph = graph - self._state_manager = state_manager - - def propagate_skip_from_edge(self, edge_id: str) -> list[GraphEdgeSkippedEvent]: - """Recursively propagate skip state from a skipped edge. - - Rules: - - If a node has any UNKNOWN incoming edges, stop processing - - If all incoming edges are SKIPPED, skip the node and its edges - - If any incoming edge is TAKEN, the node may still execute - - Args: - edge_id: The ID of the skipped edge to start from - - Returns: - Traversal events for edges marked skipped during propagation. - - """ - downstream_node_id = self._graph.edges[edge_id].head - incoming_edges = self._graph.get_incoming_edges(downstream_node_id) - - # Analyze edge states - edge_states = self._state_manager.analyze_edge_states(incoming_edges) - - # Stop if there are unknown edges (not yet processed) - if edge_states["has_unknown"]: - return [] - - # If any edge is taken, node may still execute - if edge_states["has_taken"]: - self._state_manager.enqueue_node(downstream_node_id) - return [] - - # All edges are skipped, propagate skip to this node - if edge_states["all_skipped"]: - return self._propagate_skip_to_node(downstream_node_id) - - return [] - - def propagate_skip_to_node(self, node_id: str) -> list[GraphEdgeSkippedEvent]: - """Mark a node and its downstream edges as skipped.""" - return self._propagate_skip_to_node(node_id) - - def _propagate_skip_to_node(self, node_id: str) -> list[GraphEdgeSkippedEvent]: - """Mark a node and all its outgoing edges as skipped. - - Args: - node_id: The ID of the node to skip - - Returns: - Traversal events for outgoing edges marked skipped. - - """ - # Mark node as skipped - self._state_manager.mark_node_skipped(node_id) - - # Mark all outgoing edges as skipped and propagate - events: list[GraphEdgeSkippedEvent] = [] - outgoing_edges = self._graph.get_outgoing_edges(node_id) - for edge in outgoing_edges: - events.extend(self._skip_edge_path(edge)) - return events - - def skip_branch_paths( - self, - unselected_edges: Sequence[Edge], - ) -> list[GraphEdgeSkippedEvent]: - """Skip all paths from unselected branch edges. - - Args: - unselected_edges: List of edges not taken by the branch - - Returns: - Traversal events for skipped branch edges and propagated skips. - - """ - events: list[GraphEdgeSkippedEvent] = [] - for edge in unselected_edges: - events.extend(self._skip_edge_path(edge)) - return events - - def _skip_edge_path(self, edge: Edge) -> list[GraphEdgeSkippedEvent]: - self._state_manager.mark_edge_skipped(edge.id) - return [ - self._build_skipped_event(edge), - *self.propagate_skip_from_edge(edge.id), - ] - - @staticmethod - def _build_skipped_event(edge: Edge) -> GraphEdgeSkippedEvent: - return GraphEdgeSkippedEvent( - edge_id=edge.id, - source_node_id=edge.tail, - target_node_id=edge.head, - source_handle=edge.source_handle, - ) diff --git a/src/graphon/graph_engine/layers/README.md b/src/graphon/graph_engine/layers/README.md deleted file mode 100644 index dacb4d4a..00000000 --- a/src/graphon/graph_engine/layers/README.md +++ /dev/null @@ -1,52 +0,0 @@ -# Layers - -Pluggable middleware for engine extensions. - -## Components - -### Layer (base) - -Abstract base class for layers. - -- `initialize()` - Receive runtime context (runtime state is bound here and always available to hooks) -- `on_graph_start()` - Execution start hook -- `on_event()` - Process all events -- `on_graph_end()` - Execution end hook - -### DebugLoggingLayer - -Comprehensive execution logging. - -- Configurable detail levels -- Tracks execution statistics -- Truncates long values - -## Usage - -```python -debug_layer = DebugLoggingLayer(level="INFO", include_outputs=True) - -engine = GraphEngine(graph) -engine.layer(debug_layer) -engine.run() -``` - -`engine.layer()` binds the read-only runtime state before execution, so -`graph_runtime_state` is always available inside layer hooks. - -## Custom Layers - -```python -class MetricsLayer(Layer): - def on_event(self, event): - if isinstance(event, NodeRunSucceededEvent): - self.metrics[event.node_id] = event.elapsed_time -``` - -## Configuration - -**DebugLoggingLayer Options:** - -- `level` - Log level (INFO, DEBUG, ERROR) -- `include_inputs/outputs` - Log data values -- `max_value_length` - Truncate long values diff --git a/src/graphon/graph_engine/layers/__init__.py b/src/graphon/graph_engine/layers/__init__.py deleted file mode 100644 index 03daa93b..00000000 --- a/src/graphon/graph_engine/layers/__init__.py +++ /dev/null @@ -1,15 +0,0 @@ -"""Layer system for GraphEngine extensibility. - -This module provides the layer infrastructure for extending GraphEngine functionality -with middleware-like components that can observe events and interact with execution. -""" - -from .base import GraphEngineLayer -from .debug_logging import DebugLoggingLayer -from .execution_limits import ExecutionLimitsLayer - -__all__ = [ - "DebugLoggingLayer", - "ExecutionLimitsLayer", - "GraphEngineLayer", -] diff --git a/src/graphon/graph_engine/layers/debug_logging.py b/src/graphon/graph_engine/layers/debug_logging.py deleted file mode 100644 index 4e14f742..00000000 --- a/src/graphon/graph_engine/layers/debug_logging.py +++ /dev/null @@ -1,300 +0,0 @@ -"""Debug logging layer for GraphEngine. - -This module provides a layer that logs all events and state changes during -graph execution for debugging purposes. -""" - -import logging -from collections.abc import Mapping -from functools import singledispatchmethod -from typing import Any, final, override - -from graphon.graph_events.base import GraphEngineEvent -from graphon.graph_events.graph import ( - GraphRunAbortedEvent, - GraphRunFailedEvent, - GraphRunPartialSucceededEvent, - GraphRunStartedEvent, - GraphRunSucceededEvent, -) -from graphon.graph_events.iteration import ( - NodeRunIterationFailedEvent, - NodeRunIterationNextEvent, - NodeRunIterationStartedEvent, - NodeRunIterationSucceededEvent, -) -from graphon.graph_events.loop import ( - NodeRunLoopFailedEvent, - NodeRunLoopNextEvent, - NodeRunLoopStartedEvent, - NodeRunLoopSucceededEvent, -) -from graphon.graph_events.node import ( - NodeRunExceptionEvent, - NodeRunFailedEvent, - NodeRunRetryEvent, - NodeRunStartedEvent, - NodeRunStreamChunkEvent, - NodeRunSucceededEvent, -) - -from .base import GraphEngineLayer - - -@final -class DebugLoggingLayer(GraphEngineLayer): - """A layer that provides comprehensive logging of GraphEngine execution. - - This layer logs all events with configurable detail levels, helping developers - debug workflow execution and understand the flow of events. - """ - - def __init__( - self, - level: str = "INFO", - include_inputs: bool = False, - include_outputs: bool = True, - include_process_data: bool = False, - logger_name: str = "GraphEngine.Debug", - max_value_length: int = 500, - ) -> None: - """Initialize the debug logging layer. - - Args: - level: Logging level (DEBUG, INFO, WARNING, ERROR) - include_inputs: Whether to log node input values - include_outputs: Whether to log node output values - include_process_data: Whether to log node process data - logger_name: Name of the logger to use - max_value_length: Maximum length of logged values (truncated if longer) - - """ - super().__init__() - self.level = level - self.include_inputs = include_inputs - self.include_outputs = include_outputs - self.include_process_data = include_process_data - self.max_value_length = max_value_length - - # Set up logger - self.logger = logging.getLogger(logger_name) - log_level = getattr(logging, level.upper(), logging.INFO) - self.logger.setLevel(log_level) - - # Track execution stats - self.node_count = 0 - self.success_count = 0 - self.failure_count = 0 - self.retry_count = 0 - - def _truncate_value(self, value: Any) -> str: - """Truncate long values for logging.""" - str_value = str(value) - if len(str_value) > self.max_value_length: - return str_value[: self.max_value_length] + "... (truncated)" - return str_value - - def _format_dict(self, data: dict[str, Any] | Mapping[str, Any]) -> str: - """Format a dictionary or mapping for logging with truncation.""" - if not data: - return "{}" - - formatted_items: list[str] = [] - for key, value in data.items(): - formatted_value = self._truncate_value(value) - formatted_items.append(f" {key}: {formatted_value}") - - return f"""{{ -{",\n".join(formatted_items)} -}}""" - - @override - def on_graph_start(self) -> None: - """Log graph execution start.""" - self.logger.info("=" * 80) - self.logger.info("🚀 GRAPH EXECUTION STARTED") - self.logger.info("=" * 80) - # Log initial state - self.logger.info("Initial State:") - - @override - def on_event(self, event: GraphEngineEvent) -> None: - """Log individual events based on their type.""" - self._dispatch_event(event) - - @singledispatchmethod - def _dispatch_event(self, event: GraphEngineEvent) -> None: - self.logger.debug("Event: %s", type(event).__name__) - - @_dispatch_event.register - def _log_graph_run_started_event(self, event: GraphRunStartedEvent) -> None: - _ = event - self.logger.debug("Graph run started event") - - @_dispatch_event.register - def _log_graph_run_succeeded_event(self, event: GraphRunSucceededEvent) -> None: - self.logger.info("✅ Graph run succeeded") - if self.include_outputs and event.outputs: - self.logger.info(" Final outputs: %s", self._format_dict(event.outputs)) - - @_dispatch_event.register - def _log_graph_run_partial_succeeded_event( - self, - event: GraphRunPartialSucceededEvent, - ) -> None: - self.logger.warning("⚠️ Graph run partially succeeded") - if event.exceptions_count > 0: - self.logger.warning(" Total exceptions: %s", event.exceptions_count) - if self.include_outputs and event.outputs: - self.logger.info(" Final outputs: %s", self._format_dict(event.outputs)) - - @_dispatch_event.register - def _log_graph_run_failed_event(self, event: GraphRunFailedEvent) -> None: - self.logger.error("❌ Graph run failed: %s", event.error) - if event.exceptions_count > 0: - self.logger.error(" Total exceptions: %s", event.exceptions_count) - - @_dispatch_event.register - def _log_graph_run_aborted_event(self, event: GraphRunAbortedEvent) -> None: - self.logger.warning("⚠️ Graph run aborted: %s", event.reason) - if event.outputs: - self.logger.info(" Partial outputs: %s", self._format_dict(event.outputs)) - - @_dispatch_event.register - def _log_node_run_retry_event(self, event: NodeRunRetryEvent) -> None: - self.retry_count += 1 - self.logger.warning( - "🔄 Node retry: %s (attempt %s)", - event.node_id, - event.retry_index, - ) - self.logger.warning(" Previous error: %s", event.error) - - @_dispatch_event.register - def _log_node_run_started_event(self, event: NodeRunStartedEvent) -> None: - self.node_count += 1 - self.logger.info( - '▶️ Node started: %s - "%s" (type: %s)', - event.node_id, - event.node_title, - event.node_type, - ) - if self.include_inputs and event.node_run_result.inputs: - self.logger.debug( - " Inputs: %s", - self._format_dict(event.node_run_result.inputs), - ) - - @_dispatch_event.register - def _log_node_run_succeeded_event(self, event: NodeRunSucceededEvent) -> None: - self.success_count += 1 - self.logger.info("✅ Node succeeded: %s", event.node_id) - if self.include_outputs and event.node_run_result.outputs: - self.logger.debug( - " Outputs: %s", - self._format_dict(event.node_run_result.outputs), - ) - if self.include_process_data and event.node_run_result.process_data: - self.logger.debug( - " Process data: %s", - self._format_dict(event.node_run_result.process_data), - ) - - @_dispatch_event.register - def _log_node_run_failed_event(self, event: NodeRunFailedEvent) -> None: - self.failure_count += 1 - self.logger.error("❌ Node failed: %s", event.node_id) - self.logger.error(" Error: %s", event.error) - if event.node_run_result.error: - self.logger.error(" Details: %s", event.node_run_result.error) - - @_dispatch_event.register - def _log_node_run_exception_event(self, event: NodeRunExceptionEvent) -> None: - self.logger.warning("⚠️ Node exception handled: %s", event.node_id) - self.logger.warning(" Error: %s", event.error) - - @_dispatch_event.register - def _log_node_run_stream_chunk_event(self, event: NodeRunStreamChunkEvent) -> None: - final_indicator = " (FINAL)" if event.is_final else "" - self.logger.debug( - "📝 Stream chunk from %s%s: %s", - event.node_id, - final_indicator, - self._truncate_value(event.chunk), - ) - - @_dispatch_event.register - def _log_iteration_started_event(self, event: NodeRunIterationStartedEvent) -> None: - self.logger.info("🔁 Iteration started: %s", event.node_id) - - @_dispatch_event.register - def _log_iteration_next_event(self, event: NodeRunIterationNextEvent) -> None: - self.logger.debug( - " Iteration next: %s (index: %s)", - event.node_id, - event.index, - ) - - @_dispatch_event.register - def _log_iteration_succeeded_event( - self, - event: NodeRunIterationSucceededEvent, - ) -> None: - self.logger.info("✅ Iteration succeeded: %s", event.node_id) - if self.include_outputs and event.outputs: - self.logger.debug(" Outputs: %s", self._format_dict(event.outputs)) - - @_dispatch_event.register - def _log_iteration_failed_event(self, event: NodeRunIterationFailedEvent) -> None: - self.logger.error("❌ Iteration failed: %s", event.node_id) - self.logger.error(" Error: %s", event.error) - - @_dispatch_event.register - def _log_loop_started_event(self, event: NodeRunLoopStartedEvent) -> None: - self.logger.info("🔄 Loop started: %s", event.node_id) - - @_dispatch_event.register - def _log_loop_next_event(self, event: NodeRunLoopNextEvent) -> None: - self.logger.debug( - " Loop iteration: %s (index: %s)", - event.node_id, - event.index, - ) - - @_dispatch_event.register - def _log_loop_succeeded_event(self, event: NodeRunLoopSucceededEvent) -> None: - self.logger.info("✅ Loop succeeded: %s", event.node_id) - if self.include_outputs and event.outputs: - self.logger.debug(" Outputs: %s", self._format_dict(event.outputs)) - - @_dispatch_event.register - def _log_loop_failed_event(self, event: NodeRunLoopFailedEvent) -> None: - self.logger.error("❌ Loop failed: %s", event.node_id) - self.logger.error(" Error: %s", event.error) - - @override - def on_graph_end(self, error: Exception | None) -> None: - """Log graph execution end with summary statistics.""" - self.logger.info("=" * 80) - - if error: - self.logger.error("🔴 GRAPH EXECUTION FAILED") - self.logger.error(" Error: %s", error) - else: - self.logger.info("🎉 GRAPH EXECUTION COMPLETED SUCCESSFULLY") - - # Log execution statistics - self.logger.info("Execution Statistics:") - self.logger.info(" Total nodes executed: %s", self.node_count) - self.logger.info(" Successful nodes: %s", self.success_count) - self.logger.info(" Failed nodes: %s", self.failure_count) - self.logger.info(" Node retries: %s", self.retry_count) - - # Log final state if available - if self.include_outputs and self.graph_runtime_state.outputs: - self.logger.info( - "Final outputs: %s", - self._format_dict(self.graph_runtime_state.outputs), - ) - - self.logger.info("=" * 80) diff --git a/src/graphon/graph_engine/manager.py b/src/graphon/graph_engine/manager.py deleted file mode 100644 index a4c0f87d..00000000 --- a/src/graphon/graph_engine/manager.py +++ /dev/null @@ -1,85 +0,0 @@ -"""GraphEngine Manager for sending control commands via Redis channel. - -This module provides a simplified interface for controlling workflow executions -using the new Redis command channel, without requiring user permission checks. -Callers must provide a Redis client dependency from outside the workflow package. -""" - -import logging -from collections.abc import Sequence -from typing import final - -from graphon.graph_engine.command_channels.redis_channel import ( - RedisChannel, - RedisClientProtocol, -) -from graphon.graph_engine.entities.commands import ( - AbortCommand, - GraphEngineCommand, - PauseCommand, - UpdateVariablesCommand, - VariableUpdate, -) - -logger = logging.getLogger(__name__) - - -@final -class GraphEngineManager: - """Manager for sending control commands to GraphEngine instances. - - This class provides a simple interface for controlling workflow executions - by sending commands through Redis channels, without user validation. - """ - - _redis_client: RedisClientProtocol - - def __init__(self, redis_client: RedisClientProtocol) -> None: - self._redis_client = redis_client - - def send_stop_command(self, task_id: str, reason: str | None = None) -> None: - """Send a stop command to a running workflow. - - Args: - task_id: The task ID of the workflow to stop - reason: Optional reason for stopping (defaults to "User requested stop") - - """ - abort_command = AbortCommand(reason=reason or "User requested stop") - self._send_command(task_id, abort_command) - - def send_pause_command(self, task_id: str, reason: str | None = None) -> None: - """Send a pause command to a running workflow.""" - pause_command = PauseCommand(reason=reason or "User requested pause") - self._send_command(task_id, pause_command) - - def send_update_variables_command( - self, - task_id: str, - updates: Sequence[VariableUpdate], - ) -> None: - """Send a command to update variables in a running workflow.""" - if not updates: - return - - update_command = UpdateVariablesCommand(updates=updates) - self._send_command(task_id, update_command) - - def _send_command(self, task_id: str, command: GraphEngineCommand) -> None: - """Send a command to the workflow-specific Redis channel.""" - if not task_id: - return - - channel_key = f"workflow:{task_id}:commands" - channel = RedisChannel(self._redis_client, channel_key) - - try: - channel.send_command(command) - except Exception: - # Silently fail if Redis is unavailable - # The legacy control mechanisms will still work - logger.exception( - "Failed to send graph engine command %s for task %s", - command.__class__.__name__, - task_id, - ) diff --git a/src/graphon/graph_engine/orchestration/__init__.py b/src/graphon/graph_engine/orchestration/__init__.py deleted file mode 100644 index 5ba800ff..00000000 --- a/src/graphon/graph_engine/orchestration/__init__.py +++ /dev/null @@ -1,9 +0,0 @@ -"""Orchestration subsystem for graph engine. - -This package coordinates the overall execution flow between -different subsystems. -""" - -from .dispatcher import Dispatcher - -__all__ = ["Dispatcher"] diff --git a/src/graphon/graph_engine/worker_management/__init__.py b/src/graphon/graph_engine/worker_management/__init__.py deleted file mode 100644 index f1b8959d..00000000 --- a/src/graphon/graph_engine/worker_management/__init__.py +++ /dev/null @@ -1,11 +0,0 @@ -"""Worker management subsystem for graph engine. - -This package manages the worker pool, including creation, -scaling, and activity tracking. -""" - -from .worker_pool import WorkerPool - -__all__ = [ - "WorkerPool", -] diff --git a/src/graphon/graph_engine/worker_management/worker_pool.py b/src/graphon/graph_engine/worker_management/worker_pool.py deleted file mode 100644 index fca62aa8..00000000 --- a/src/graphon/graph_engine/worker_management/worker_pool.py +++ /dev/null @@ -1,278 +0,0 @@ -"""Simple worker pool that consolidates functionality. - -This is a simpler implementation that merges WorkerPool, ActivityTracker, -DynamicScaler, and WorkerFactory into a single class. -""" - -import logging -import queue -import threading -from contextlib import AbstractContextManager -from typing import final - -from graphon.graph_engine.entities.tasks import DispatchTask -from graphon.graph_engine.frames import FrameRegistry -from graphon.graph_engine.ready_queue import ROOT_FRAME_ID, ReadyQueue, ReadyTask - -from ..config import GraphEngineConfig -from ..layers.base import GraphEngineLayer -from ..worker import Worker - -logger = logging.getLogger(__name__) -SMALL_GRAPH_NODE_THRESHOLD = 10 -MEDIUM_GRAPH_NODE_THRESHOLD = 50 - - -@final -class WorkerPool: - """Simple worker pool with integrated management. - - This class consolidates all worker management functionality into - a single, simpler implementation without excessive abstraction. - """ - - def __init__( - self, - ready_queue: ReadyQueue, - event_queue: queue.Queue[DispatchTask], - frame_registry: FrameRegistry, - layers: list[GraphEngineLayer], - config: GraphEngineConfig, - execution_context: AbstractContextManager[object] | None = None, - ) -> None: - """Initialize the simple worker pool. - - Args: - ready_queue: Ready queue protocol for nodes ready for execution - event_queue: Queue for worker events - frame_registry: Registry containing frame-local graphs to execute - layers: Graph engine layers for node execution hooks - config: GraphEngine worker pool configuration - execution_context: Optional execution context for context preservation - - """ - self._ready_queue = ready_queue - self._event_queue = event_queue - self._frame_registry = frame_registry - self._execution_context = execution_context - self._layers = layers - self._config = config - - # Worker management - self._workers: list[Worker] = [] - self._worker_counter = 0 - self._lock = threading.Lock() - self._task_claim_lock = threading.Lock() - self._task_claiming = threading.Event() - self._running = False - - def start(self) -> None: - """Start the worker pool.""" - with self._lock: - if self._running: - return - - self._running = True - self._task_claiming.set() - node_count = len( - self._frame_registry.get(ROOT_FRAME_ID).graph.nodes, - ) - if node_count < SMALL_GRAPH_NODE_THRESHOLD: - initial_count = self._config.min_workers - elif node_count < MEDIUM_GRAPH_NODE_THRESHOLD: - initial_count = min( - self._config.min_workers + 1, - self._config.max_workers, - ) - else: - initial_count = min( - self._config.min_workers + 2, - self._config.max_workers, - ) - - logger.debug( - "Starting worker pool: %d workers (nodes=%d, min=%d, max=%d)", - initial_count, - node_count, - self._config.min_workers, - self._config.max_workers, - ) - for _ in range(initial_count): - self._create_worker() - - def stop(self) -> None: - """Stop all workers in the pool.""" - with self._lock: - self._running = False - with self._task_claim_lock: - self._task_claiming.clear() - worker_count = len(self._workers) - - if worker_count > 0: - logger.debug("Stopping worker pool: %d workers", worker_count) - - # Stop all workers - for worker in self._workers: - worker.stop() - - # Wait for workers to finish - for worker in self._workers: - if worker.is_alive(): - worker.join(timeout=2.0) - - self._workers.clear() - - def drain(self) -> list[ReadyTask]: - """Atomically stop task claims and remove unclaimed ready work.""" - with self._lock: - self._running = False - with self._task_claim_lock: - self._task_claiming.clear() - tasks = self._ready_queue.drain() - for worker in self._workers: - if not worker.has_current_task: - worker.stop() - return tasks - - def has_current_tasks(self) -> bool: - with self._lock: - return any(worker.has_current_task for worker in self._workers) - - def _create_worker(self) -> None: - """Create and start a new worker.""" - worker_id = self._worker_counter - self._worker_counter += 1 - - worker = Worker( - ready_queue=self._ready_queue, - event_queue=self._event_queue, - frame_registry=self._frame_registry, - layers=self._layers, - worker_id=worker_id, - execution_context=self._execution_context, - task_claim_lock=self._task_claim_lock, - task_claiming=self._task_claiming, - ) - - worker.start() - self._workers.append(worker) - - def _remove_worker(self, worker: Worker) -> None: - """Remove a specific worker from the pool.""" - # Stop the worker - worker.stop() - - # Wait for it to finish - if worker.is_alive(): - worker.join(timeout=2.0) - - self._workers.remove(worker) - - def _try_scale_up( - self, - queue_depth: int, - current_count: int, - active_count: int, - ) -> bool: - """Try to scale up workers if needed. - - Args: - queue_depth: Current queue depth - current_count: Current number of workers - - Returns: - True if scaled up, False otherwise - - """ - available_count = current_count - active_count - backlog = max(queue_depth - available_count, 0) - if backlog > self._config.scale_up_threshold and ( - current_count < self._config.max_workers - ): - self._create_worker() - - logger.debug( - "Scaled up workers: %d -> %d (backlog=%d exceeded threshold=%d)", - current_count, - len(self._workers), - backlog, - self._config.scale_up_threshold, - ) - return True - return False - - def _try_scale_down( - self, - queue_depth: int, - current_count: int, - active_count: int, - idle_count: int, - ) -> bool: - """Try to scale down workers if we have excess capacity. - - Args: - queue_depth: Current queue depth - current_count: Current number of workers - active_count: Number of active workers - idle_count: Number of idle workers - - Returns: - True if scaled down, False otherwise - - """ - # Skip if we're at minimum or have no idle workers - if current_count <= self._config.min_workers or idle_count == 0: - return False - - # Check if we have excess capacity - has_excess_capacity = ( - queue_depth <= active_count # Active workers can handle current queue - or idle_count > active_count # More idle than active workers - ) - - if not has_excess_capacity: - return False - - for worker in self._workers: - if ( - worker.is_idle - and worker.idle_duration >= self._config.scale_down_idle_time - ): - remaining_workers = current_count - 1 - if ( - remaining_workers >= self._config.min_workers - and remaining_workers >= max(1, queue_depth // 2) - ): - self._remove_worker(worker) - logger.debug( - "Scaled down workers: %d -> %d (removed 1 idle worker after " - "%.1fs, queue_depth=%d, active=%d, idle=%d)", - current_count, - len(self._workers), - self._config.scale_down_idle_time, - queue_depth, - active_count, - idle_count - 1, - ) - return True - - return False - - def check_and_scale(self) -> None: - """Check and perform scaling if needed.""" - with self._lock: - if not self._running: - return - - current_count = len(self._workers) - queue_depth = self._ready_queue.qsize() - - # Active ownership is immediate; idle status includes the scale-down delay. - active_count = sum(1 for worker in self._workers if worker.has_current_task) - idle_count = sum(1 for worker in self._workers if worker.is_idle) - - # Try to scale up if queue is backing up - self._try_scale_up(queue_depth, current_count, active_count) - - # Try to scale down if we have excess capacity - self._try_scale_down(queue_depth, current_count, active_count, idle_count) diff --git a/src/graphon/graph_events/human_input.py b/src/graphon/graph_events/human_input.py deleted file mode 100644 index e69de29b..00000000 diff --git a/src/graphon/node_events/__init__.py b/src/graphon/node_events/__init__.py index beca2e5e..fe896a1d 100644 --- a/src/graphon/node_events/__init__.py +++ b/src/graphon/node_events/__init__.py @@ -1,5 +1,5 @@ from .agent import AgentLogEvent -from .base import NodeEventBase, NodeRunResult +from .base import NodeEventPayload, NodeRunResult from .iteration import ( IterationFailedEvent, IterationNextEvent, @@ -40,7 +40,7 @@ "LoopSucceededEvent", "ModelInvokeCompletedEvent", "ModelPollingProgressEvent", - "NodeEventBase", + "NodeEventPayload", "NodeRunResult", "PauseRequestedEvent", "RunRetrieverResourceEvent", diff --git a/src/graphon/node_events/agent.py b/src/graphon/node_events/agent.py index bf295ec7..02d5e0a2 100644 --- a/src/graphon/node_events/agent.py +++ b/src/graphon/node_events/agent.py @@ -3,10 +3,10 @@ from pydantic import Field -from .base import NodeEventBase +from .base import NodeEventPayload -class AgentLogEvent(NodeEventBase): +class AgentLogEvent(NodeEventPayload): message_id: str = Field(..., description="id") label: str = Field(..., description="label") node_execution_id: str = Field(..., description="node execution id") diff --git a/src/graphon/node_events/base.py b/src/graphon/node_events/base.py index 6d5b78fb..f57abaac 100644 --- a/src/graphon/node_events/base.py +++ b/src/graphon/node_events/base.py @@ -7,8 +7,8 @@ from graphon.model_runtime.entities.llm_entities import LLMUsage -class NodeEventBase(BaseModel): - """Base class for all node events""" +class NodeEventPayload(BaseModel): + """Event payload emitted by a node before execution context is attached.""" def _default_metadata() -> Mapping[WorkflowNodeExecutionMetadataKey, Any]: diff --git a/src/graphon/node_events/iteration.py b/src/graphon/node_events/iteration.py index bbcfd348..f4c437b8 100644 --- a/src/graphon/node_events/iteration.py +++ b/src/graphon/node_events/iteration.py @@ -3,22 +3,22 @@ from pydantic import Field -from .base import NodeEventBase +from .base import NodeEventPayload -class IterationStartedEvent(NodeEventBase): +class IterationStartedEvent(NodeEventPayload): start_at: datetime = Field(..., description="start at") inputs: Mapping[str, object] = Field(default_factory=dict) metadata: Mapping[str, object] = Field(default_factory=dict) predecessor_node_id: str | None = None -class IterationNextEvent(NodeEventBase): +class IterationNextEvent(NodeEventPayload): index: int = Field(..., description="index") pre_iteration_output: object = None -class IterationSucceededEvent(NodeEventBase): +class IterationSucceededEvent(NodeEventPayload): start_at: datetime = Field(..., description="start at") inputs: Mapping[str, object] = Field(default_factory=dict) outputs: Mapping[str, object] = Field(default_factory=dict) @@ -26,7 +26,7 @@ class IterationSucceededEvent(NodeEventBase): steps: int = 0 -class IterationFailedEvent(NodeEventBase): +class IterationFailedEvent(NodeEventPayload): start_at: datetime = Field(..., description="start at") inputs: Mapping[str, object] = Field(default_factory=dict) outputs: Mapping[str, object] = Field(default_factory=dict) diff --git a/src/graphon/node_events/loop.py b/src/graphon/node_events/loop.py index 08e0e6e5..3a53ebc2 100644 --- a/src/graphon/node_events/loop.py +++ b/src/graphon/node_events/loop.py @@ -3,22 +3,22 @@ from pydantic import Field -from .base import NodeEventBase +from .base import NodeEventPayload -class LoopStartedEvent(NodeEventBase): +class LoopStartedEvent(NodeEventPayload): start_at: datetime = Field(..., description="start at") inputs: Mapping[str, object] = Field(default_factory=dict) metadata: Mapping[str, object] = Field(default_factory=dict) predecessor_node_id: str | None = None -class LoopNextEvent(NodeEventBase): +class LoopNextEvent(NodeEventPayload): index: int = Field(..., description="index") pre_loop_output: object = None -class LoopSucceededEvent(NodeEventBase): +class LoopSucceededEvent(NodeEventPayload): start_at: datetime = Field(..., description="start at") inputs: Mapping[str, object] = Field(default_factory=dict) outputs: Mapping[str, object] = Field(default_factory=dict) @@ -26,7 +26,7 @@ class LoopSucceededEvent(NodeEventBase): steps: int = 0 -class LoopFailedEvent(NodeEventBase): +class LoopFailedEvent(NodeEventPayload): start_at: datetime = Field(..., description="start at") inputs: Mapping[str, object] = Field(default_factory=dict) outputs: Mapping[str, object] = Field(default_factory=dict) diff --git a/src/graphon/node_events/node.py b/src/graphon/node_events/node.py index 7bbb5c89..aac2c62b 100644 --- a/src/graphon/node_events/node.py +++ b/src/graphon/node_events/node.py @@ -11,10 +11,10 @@ from graphon.variables.segments import Segment from graphon.variables.variables import Variable -from .base import NodeEventBase +from .base import NodeEventPayload -class RunRetrieverResourceEvent(NodeEventBase): +class RunRetrieverResourceEvent(NodeEventPayload): retriever_resources: Sequence[Mapping[str, Any]] = Field( ..., description="retriever resources", @@ -23,7 +23,7 @@ class RunRetrieverResourceEvent(NodeEventBase): context_files: list[File] | None = Field(default=None, description="context files") -class ModelInvokeCompletedEvent(NodeEventBase): +class ModelInvokeCompletedEvent(NodeEventPayload): text: str usage: LLMUsage finish_reason: str | None = None @@ -31,7 +31,7 @@ class ModelInvokeCompletedEvent(NodeEventBase): structured_output: dict | None = None -class ModelPollingProgressEvent(NodeEventBase): +class ModelPollingProgressEvent(NodeEventPayload): attempt: int = Field(..., ge=0, description="polling check attempt count") last_checked_at: datetime = Field(..., description="last polling check time") next_check_at: datetime | None = Field( @@ -40,13 +40,13 @@ class ModelPollingProgressEvent(NodeEventBase): ) -class RunRetryEvent(NodeEventBase): +class RunRetryEvent(NodeEventPayload): error: str = Field(..., description="error") retry_index: int = Field(..., description="Retry attempt number") start_at: datetime = Field(..., description="Retry start time") -class StreamChunkEvent(NodeEventBase): +class StreamChunkEvent(NodeEventPayload): # Spec-compliant fields selector: Sequence[str] = Field( ..., @@ -61,7 +61,7 @@ class StreamChunkEvent(NodeEventBase): ) -class StreamReasoningEvent(NodeEventBase): +class StreamReasoningEvent(NodeEventPayload): """Reasoning side-channel chunk scoped by selector visibility. This node-level selector is normalized by graph dispatch to @@ -86,21 +86,21 @@ class StreamReasoningEvent(NodeEventBase): ) -class StreamCompletedEvent(NodeEventBase): +class StreamCompletedEvent(NodeEventPayload): node_run_result: NodeRunResult = Field(..., description="run result") -class VariableUpdatedEvent(NodeEventBase): +class VariableUpdatedEvent(NodeEventPayload): """Notify the engine that a single variable should be applied to the shared pool.""" variable: Variable = Field(..., description="Updated variable payload to apply.") -class PauseRequestedEvent(NodeEventBase): +class PauseRequestedEvent(NodeEventPayload): reason: PauseReason = Field(..., description="pause reason") -class HumanInputFormFilledEvent(NodeEventBase): +class HumanInputFormFilledEvent(NodeEventPayload): """Event emitted when a human input form is submitted.""" node_title: str @@ -114,7 +114,7 @@ class HumanInputFormFilledEvent(NodeEventBase): submitted_data: Mapping[str, Segment] = Field(default_factory=dict) -class HumanInputFormTimeoutEvent(NodeEventBase): +class HumanInputFormTimeoutEvent(NodeEventPayload): """Event emitted when a human input form times out.""" node_title: str diff --git a/src/graphon/nodes/base/node.py b/src/graphon/nodes/base/node.py index 4da72399..dc70dead 100644 --- a/src/graphon/nodes/base/node.py +++ b/src/graphon/nodes/base/node.py @@ -9,31 +9,21 @@ from types import MappingProxyType from typing import Any, ClassVar, assert_never, get_args, get_origin -from graphon.entities.base_node_data import BaseNodeData, RetryConfig -from graphon.entities.graph_config import NodeConfigDict, NodeConfigDictAdapter -from graphon.entities.graph_init_params import GraphInitParams -from graphon.enums import ( - ErrorStrategy, - NodeExecutionType, - NodeState, - NodeType, - WorkflowNodeExecutionStatus, -) -from graphon.graph_events.agent import NodeRunAgentLogEvent -from graphon.graph_events.base import GraphNodeEventBase -from graphon.graph_events.iteration import ( +from graphon.engine_events.agent import NodeRunAgentLogEvent +from graphon.engine_events.base import NodeEvent +from graphon.engine_events.iteration import ( NodeRunIterationFailedEvent, NodeRunIterationNextEvent, NodeRunIterationStartedEvent, NodeRunIterationSucceededEvent, ) -from graphon.graph_events.loop import ( +from graphon.engine_events.loop import ( NodeRunLoopFailedEvent, NodeRunLoopNextEvent, NodeRunLoopStartedEvent, NodeRunLoopSucceededEvent, ) -from graphon.graph_events.node import ( +from graphon.engine_events.node import ( NodeRunFailedEvent, NodeRunHumanInputFormFilledEvent, NodeRunHumanInputFormTimeoutEvent, @@ -46,9 +36,19 @@ NodeRunSucceededEvent, NodeRunVariableUpdatedEvent, ) +from graphon.entities.base_node_data import BaseNodeData, RetryConfig +from graphon.entities.graph_config import NodeConfigDict, NodeConfigDictAdapter +from graphon.entities.graph_init_params import InitParams +from graphon.enums import ( + ErrorStrategy, + NodeExecutionType, + NodeState, + NodeType, + WorkflowNodeExecutionStatus, +) from graphon.node_events.agent import AgentLogEvent from graphon.node_events.base import ( - NodeEventBase, + NodeEventPayload, NodeRunResult, ) from graphon.node_events.iteration import ( @@ -75,7 +75,7 @@ VariableUpdatedEvent, ) from graphon.nodes.container_effects import ContainerAwaitRequest, ContainerRunResult -from graphon.runtime.graph_runtime_state import GraphRuntimeState +from graphon.runtime.graph_runtime_state import RuntimeState _MISSING_RUN_CONTEXT_VALUE = object() @@ -335,8 +335,26 @@ def post_init(self: Node[NodeDataT]) -> None: """Optional hook for subclasses requiring extra initialization.""" return + def bind_graph_config( + self: Node[NodeDataT], + graph_config: Mapping[str, Any], + ) -> None: + """Limit graph configuration visible to this node's execution frame. + + Both public access paths are updated: ``node.graph_config`` and + ``node.graph_init_params.graph_config``. A private copy of + :class:`InitParams` prevents rebinding this node from changing the + factory-owned params shared by nodes in parent or sibling frames. + + :param graph_config: current frame's container-subtree configuration + """ + self.graph_config = graph_config + self._graph_init_params = self._graph_init_params.model_copy( + update={"graph_config": graph_config}, + ) + @property - def graph_init_params(self: Node[NodeDataT]) -> GraphInitParams: + def graph_init_params(self: Node[NodeDataT]) -> InitParams: return self._graph_init_params @property @@ -593,8 +611,8 @@ def __init__( node_id: str, data: NodeDataT, *, - graph_init_params: GraphInitParams, - graph_runtime_state: GraphRuntimeState, + graph_init_params: InitParams, + graph_runtime_state: RuntimeState, ) -> None: if not node_id: msg = "node_id is required" @@ -623,7 +641,7 @@ def _run( ) -> ( NodeRunResult | Generator[ - NodeEventBase | GraphNodeEventBase | ContainerAwaitRequest, + NodeEventPayload | NodeEvent | ContainerAwaitRequest, None, None, ] @@ -634,7 +652,7 @@ def _run( def run( self, ) -> Generator[ - GraphNodeEventBase | ContainerAwaitRequest, + NodeEvent | ContainerAwaitRequest, None, None, ]: @@ -647,7 +665,6 @@ def run( node_id=self._node_id, node_type=self.node_type, node_title=self.title, - in_iteration_id=None, start_at=self._start_at, ) try: @@ -669,13 +686,13 @@ def run( def _run_events( self, ) -> Generator[ - GraphNodeEventBase | ContainerAwaitRequest, + NodeEvent | ContainerAwaitRequest, None, None, ]: result = self._run() if isinstance(result, NodeRunResult): - yield self._convert_node_run_result_to_graph_node_event(result) + yield self._convert_node_run_result_to_node_event(result) return for event in result: @@ -686,7 +703,7 @@ def resume_container( *, result: ContainerRunResult, started_at: datetime, - ) -> Generator[GraphNodeEventBase | ContainerAwaitRequest, None, None]: + ) -> Generator[NodeEvent | ContainerAwaitRequest, None, None]: self._start_at = started_at try: for event in self._resume_container_events(result=result): @@ -700,7 +717,7 @@ def _resume_container_events( *, result: ContainerRunResult, ) -> Generator[ - NodeEventBase | GraphNodeEventBase | ContainerAwaitRequest, + NodeEventPayload | NodeEvent | ContainerAwaitRequest, None, None, ]: @@ -710,13 +727,13 @@ def _resume_container_events( def _normalize_run_event( self, - event: NodeEventBase | GraphNodeEventBase | ContainerAwaitRequest, - ) -> GraphNodeEventBase | ContainerAwaitRequest: + event: NodeEventPayload | NodeEvent | ContainerAwaitRequest, + ) -> NodeEvent | ContainerAwaitRequest: if isinstance(event, ContainerAwaitRequest): return event - if isinstance(event, NodeEventBase): + if isinstance(event, NodeEventPayload): return self._dispatch(event) - if not event.in_iteration_id and not event.in_loop_id: + if not event.container_id: event.id = self.execution_id return event @@ -816,10 +833,10 @@ def version(cls) -> str: msg = "subclasses of BaseNode must implement `version` method." raise NotImplementedError(msg) - def _convert_node_run_result_to_graph_node_event( + def _convert_node_run_result_to_node_event( self, result: NodeRunResult, - ) -> GraphNodeEventBase: + ) -> NodeEvent: finished_at = datetime.now(UTC).replace(tzinfo=None) status = result.status match status: @@ -856,7 +873,7 @@ def _convert_node_run_result_to_graph_node_event( assert_never(status) @singledispatchmethod - def _dispatch(self, event: NodeEventBase) -> GraphNodeEventBase: + def _dispatch(self, event: NodeEventPayload) -> NodeEvent: msg = f"Node {self._node_id} does not support event type {type(event)}" raise NotImplementedError(msg) diff --git a/src/graphon/nodes/code/code_node.py b/src/graphon/nodes/code/code_node.py index fbdb2e68..2f0f2db3 100644 --- a/src/graphon/nodes/code/code_node.py +++ b/src/graphon/nodes/code/code_node.py @@ -6,13 +6,13 @@ from textwrap import dedent from typing import Any, Protocol, TypeGuard, assert_never, override -from graphon.entities.graph_init_params import GraphInitParams +from graphon.entities.graph_init_params import InitParams from graphon.enums import BuiltinNodeTypes, WorkflowNodeExecutionStatus from graphon.node_events.base import NodeRunResult from graphon.nodes.base.node import Node from graphon.nodes.code.entities import CodeLanguage, CodeNodeData from graphon.nodes.code.limits import CodeNodeLimits -from graphon.runtime.graph_runtime_state import GraphRuntimeState +from graphon.runtime.graph_runtime_state import RuntimeState from graphon.variables.segments import ArrayFileSegment from graphon.variables.types import SegmentType @@ -87,8 +87,8 @@ def __init__( node_id: str, data: CodeNodeData, *, - graph_init_params: GraphInitParams, - graph_runtime_state: GraphRuntimeState, + graph_init_params: InitParams, + graph_runtime_state: RuntimeState, code_executor: CodeExecutorProtocol, code_limits: CodeNodeLimits, ) -> None: diff --git a/src/graphon/nodes/container_effects.py b/src/graphon/nodes/container_effects.py index 98e3bc07..e72c8d44 100644 --- a/src/graphon/nodes/container_effects.py +++ b/src/graphon/nodes/container_effects.py @@ -72,7 +72,16 @@ class IterationFrameRequest(BaseModel): parallel_nums: int -ContainerAwaitRequest = LoopFrameRequest | IterationFrameRequest +class CustomContainerRequest(BaseModel): + model_config = ConfigDict(frozen=True) + + kind: Literal["custom"] = "custom" + payload: str + + +ContainerAwaitRequest = ( + LoopFrameRequest | IterationFrameRequest | CustomContainerRequest +) class ContainerExecutionResult(BaseModel): diff --git a/src/graphon/nodes/document_extractor/node.py b/src/graphon/nodes/document_extractor/node.py index 7478dd06..b7813aa5 100644 --- a/src/graphon/nodes/document_extractor/node.py +++ b/src/graphon/nodes/document_extractor/node.py @@ -29,7 +29,7 @@ from docx.text.paragraph import Paragraph from odfdo import Document as OdfDocument -from graphon.entities.graph_init_params import GraphInitParams +from graphon.entities.graph_init_params import InitParams from graphon.enums import BuiltinNodeTypes, WorkflowNodeExecutionStatus from graphon.file import file_manager from graphon.file.enums import FileTransferMethod @@ -37,7 +37,7 @@ from graphon.http import HttpClientProtocol, get_http_client from graphon.node_events.base import NodeRunResult from graphon.nodes.base.node import Node -from graphon.runtime.graph_runtime_state import GraphRuntimeState +from graphon.runtime.graph_runtime_state import RuntimeState from graphon.variables.segments import ArrayFileSegment, ArrayStringSegment, FileSegment from .entities import DocumentExtractorNodeData, UnstructuredApiConfig @@ -288,8 +288,8 @@ def __init__( node_id: str, data: DocumentExtractorNodeData, *, - graph_init_params: GraphInitParams, - graph_runtime_state: GraphRuntimeState, + graph_init_params: InitParams, + graph_runtime_state: RuntimeState, unstructured_api_config: UnstructuredApiConfig | None = None, http_client: HttpClientProtocol | None = None, ) -> None: diff --git a/src/graphon/nodes/http_request/node.py b/src/graphon/nodes/http_request/node.py index 461f4789..34570cca 100644 --- a/src/graphon/nodes/http_request/node.py +++ b/src/graphon/nodes/http_request/node.py @@ -6,7 +6,7 @@ from dataclasses import dataclass from typing import Any, assert_never, cast, override -from graphon.entities.graph_init_params import GraphInitParams +from graphon.entities.graph_init_params import InitParams from graphon.enums import BuiltinNodeTypes, WorkflowNodeExecutionStatus from graphon.file.enums import FileTransferMethod from graphon.file.models import File @@ -21,7 +21,7 @@ FileReferenceFactoryProtocol, ToolFileManagerProtocol, ) -from graphon.runtime.graph_runtime_state import GraphRuntimeState +from graphon.runtime.graph_runtime_state import RuntimeState from graphon.variables.segments import ArrayFileSegment from .config import build_http_request_config, resolve_http_request_config @@ -57,8 +57,8 @@ def __init__( node_id: str, data: HttpRequestNodeData, *, - graph_init_params: GraphInitParams, - graph_runtime_state: GraphRuntimeState, + graph_init_params: InitParams, + graph_runtime_state: RuntimeState, http_request_config: HttpRequestNodeConfig, dependencies: HttpRequestNodeDependencies | None = None, http_client: HttpClientProtocol | None = None, diff --git a/src/graphon/nodes/human_input/human_input_node.py b/src/graphon/nodes/human_input/human_input_node.py index 50df6549..eee0b49d 100644 --- a/src/graphon/nodes/human_input/human_input_node.py +++ b/src/graphon/nodes/human_input/human_input_node.py @@ -3,17 +3,17 @@ from collections.abc import Generator, Mapping, Sequence from typing import Any, override -from graphon.entities.graph_init_params import GraphInitParams +from graphon.entities.graph_init_params import InitParams from graphon.entities.pause_reason import HitlRequired from graphon.enums import ( BuiltinNodeTypes, NodeExecutionType, WorkflowNodeExecutionStatus, ) -from graphon.node_events.base import NodeEventBase, NodeRunResult +from graphon.node_events.base import NodeEventPayload, NodeRunResult from graphon.node_events.node import PauseRequestedEvent, StreamCompletedEvent from graphon.nodes.base.node import Node -from graphon.runtime.graph_runtime_state import GraphRuntimeState +from graphon.runtime.graph_runtime_state import RuntimeState from graphon.variables.segments import Segment from .entities import ( @@ -42,8 +42,8 @@ def __init__( node_id: str, data: HumanInputNodeData, *, - graph_init_params: GraphInitParams, - graph_runtime_state: GraphRuntimeState, + graph_init_params: InitParams, + graph_runtime_state: RuntimeState, hitl_callback: HITLCallback, ) -> None: super().__init__( @@ -60,7 +60,7 @@ def version(cls) -> str: return "1" @override - def _run(self) -> Generator[NodeEventBase, None, None]: + def _run(self) -> Generator[NodeEventPayload, None, None]: decision = self._hitl_callback( HITLContext( workflow_execution_id=self._resolve_workflow_execution_id(), diff --git a/src/graphon/nodes/iteration/iteration_node.py b/src/graphon/nodes/iteration/iteration_node.py index c0b43d44..5cc0f820 100644 --- a/src/graphon/nodes/iteration/iteration_node.py +++ b/src/graphon/nodes/iteration/iteration_node.py @@ -8,7 +8,7 @@ NodeExecutionType, WorkflowNodeExecutionStatus, ) -from graphon.node_events.base import NodeEventBase, NodeRunResult +from graphon.node_events.base import NodeEventPayload, NodeRunResult from graphon.node_events.iteration import ( IterationFailedEvent, IterationNextEvent, @@ -30,7 +30,7 @@ class IterationNode(Node[IterationNodeData]): """Iteration node definition. - Iteration execution is interpreted by GraphEngine. The node keeps only its + Iteration execution is interpreted by Engine. The node keeps only its configuration and static variable-mapping behavior. """ @@ -62,7 +62,7 @@ def version(cls) -> str: @override def _run( self, - ) -> Generator[NodeEventBase | IterationFrameRequest, None, None]: + ) -> Generator[NodeEventPayload | IterationFrameRequest, None, None]: variable = self.graph_runtime_state.variable_pool.get( self.node_data.iterator_selector, ) @@ -108,7 +108,7 @@ def _resume_container_events( self, *, result: ContainerRunResult, - ) -> Generator[NodeEventBase | IterationFrameRequest, None, None]: + ) -> Generator[NodeEventPayload | IterationFrameRequest, None, None]: if isinstance(result, IterationFrameRequest): for index in result.indexes: yield IterationNextEvent(index=index) @@ -158,7 +158,7 @@ def _run_empty_iteration( *, variable: NoneSegment | ArraySegment, started_at: datetime, - ) -> Generator[NodeEventBase, None, None]: + ) -> Generator[NodeEventPayload, None, None]: outputs = {"output": ArrayAnySegment(value=[])} if isinstance(variable, ArraySegment): outputs = {"output": variable.model_copy(update={"value": []})} diff --git a/src/graphon/nodes/llm/node.py b/src/graphon/nodes/llm/node.py index e69f1578..67815899 100644 --- a/src/graphon/nodes/llm/node.py +++ b/src/graphon/nodes/llm/node.py @@ -15,7 +15,7 @@ from referencing import Registry from referencing.exceptions import Unresolvable -from graphon.entities.graph_init_params import GraphInitParams +from graphon.entities.graph_init_params import InitParams from graphon.enums import ( BuiltinNodeTypes, WorkflowNodeExecutionMetadataKey, @@ -47,7 +47,7 @@ from graphon.model_runtime.memory.prompt_message_memory import PromptMessageMemory from graphon.model_runtime.utils.encoders import jsonable_encoder from graphon.node_events.base import ( - NodeEventBase, + NodeEventPayload, NodeRunResult, ) from graphon.node_events.node import ( @@ -74,7 +74,7 @@ RetrieverAttachmentLoaderProtocol, ) from graphon.prompt_entities import MemoryConfig -from graphon.runtime.graph_runtime_state import GraphRuntimeState +from graphon.runtime.graph_runtime_state import RuntimeState from graphon.template_rendering import Jinja2TemplateRenderer from graphon.variables.segments import ( ArrayFileSegment, @@ -154,8 +154,8 @@ def __init__( node_id: str, data: LLMNodeData, *, - graph_init_params: GraphInitParams, - graph_runtime_state: GraphRuntimeState, + graph_init_params: InitParams, + graph_runtime_state: RuntimeState, credentials_provider: object | None = None, model_factory: object | None = None, model_instance: LLMProtocol, @@ -243,7 +243,7 @@ def _prepare_run_prompt( *, node_inputs: dict[str, Any], ) -> Generator[ - NodeEventBase, + NodeEventPayload, None, _PreparedRunPrompt, ]: @@ -309,7 +309,7 @@ def _collect_run_context( self, *, node_inputs: dict[str, Any], - ) -> Generator[NodeEventBase, None, _CollectedRunContext]: + ) -> Generator[NodeEventPayload, None, _CollectedRunContext]: context = None context_files: Sequence[File] = () for event in self._fetch_context(node_data=self.node_data): @@ -362,7 +362,7 @@ def _yield_run_completion( stop: Sequence[str] | None, model_provider: Any, model_name: str, - ) -> Generator[NodeEventBase, None, None]: + ) -> Generator[NodeEventPayload, None, None]: generator = self._invoke_llm_for_run( file_outputs=file_outputs, prompt_messages=prompt_messages, @@ -450,7 +450,7 @@ def _invoke_llm_for_run( file_outputs: list[File], prompt_messages: Sequence[PromptMessage], stop: Sequence[str] | None, - ) -> Generator[NodeEventBase | LLMStructuredOutput, None, None]: + ) -> Generator[NodeEventPayload | LLMStructuredOutput, None, None]: polling_model = self._polling_model_instance() if polling_model is None: return LLMNode.invoke_llm( @@ -484,7 +484,7 @@ def _invoke_llm_with_polling( polling_model: LLMPollingCapableProtocol, prompt_messages: Sequence[PromptMessage], stop: Sequence[str] | None, - ) -> Generator[NodeEventBase | LLMStructuredOutput, None, None]: + ) -> Generator[NodeEventPayload | LLMStructuredOutput, None, None]: config = self._polling_config(polling_model) model_parameters = dict(self._model_instance.parameters) json_schema = ( @@ -741,7 +741,7 @@ def invoke_llm( file_outputs: list[File], node_id: str, reasoning_format: Literal["separated", "tagged"] = "tagged", - ) -> Generator[NodeEventBase | LLMStructuredOutput, None, None]: + ) -> Generator[NodeEventPayload | LLMStructuredOutput, None, None]: model_parameters = model_instance.parameters invoke_model_parameters = dict(model_parameters) invoke_result: LLMResult | Generator[LLMResultChunk, None, None] @@ -793,7 +793,7 @@ def handle_invoke_result( reasoning_format: Literal["separated", "tagged"] = "tagged", request_start_time: float | None = None, json_schema: Mapping[str, Any] | None = None, - ) -> Generator[NodeEventBase | LLMStructuredOutput, None, None]: + ) -> Generator[NodeEventPayload | LLMStructuredOutput, None, None]: if isinstance(invoke_result, LLMResult): yield from LLMNode._yield_blocking_invoke_result( invoke_result=invoke_result, @@ -855,7 +855,7 @@ def _yield_streaming_invoke_result( reasoning_format: Literal["separated", "tagged"] = "tagged", request_start_time: float | None = None, json_schema: Mapping[str, Any] | None = None, - ) -> Generator[NodeEventBase | LLMStructuredOutput, None, None]: + ) -> Generator[NodeEventPayload | LLMStructuredOutput, None, None]: start_time = ( request_start_time if request_start_time is not None @@ -943,7 +943,7 @@ def _yield_streaming_events( file_saver: LLMFileSaver, file_outputs: list[File], node_id: str, - ) -> Generator[NodeEventBase | LLMStructuredOutput, None, None]: + ) -> Generator[NodeEventPayload | LLMStructuredOutput, None, None]: for result in invoke_result: yield from LLMNode._handle_stream_result( result=result, @@ -961,7 +961,7 @@ def _handle_stream_result( file_saver: LLMFileSaver, file_outputs: list[File], node_id: str, - ) -> Generator[NodeEventBase | LLMStructuredOutput, None, None]: + ) -> Generator[NodeEventPayload | LLMStructuredOutput, None, None]: if isinstance(result, LLMResultChunkWithStructuredOutput): if result.structured_output is not None: state.structured_output = dict(result.structured_output) diff --git a/src/graphon/nodes/loop/loop_node.py b/src/graphon/nodes/loop/loop_node.py index 64fe8517..c9e29180 100644 --- a/src/graphon/nodes/loop/loop_node.py +++ b/src/graphon/nodes/loop/loop_node.py @@ -9,7 +9,7 @@ NodeExecutionType, WorkflowNodeExecutionStatus, ) -from graphon.node_events.base import NodeEventBase +from graphon.node_events.base import NodeEventPayload from graphon.node_events.loop import ( LoopFailedEvent, LoopNextEvent, @@ -39,7 +39,7 @@ class LoopNode(Node[LoopNodeData]): """Loop node definition. - Loop execution is interpreted by GraphEngine. The node keeps only its + Loop execution is interpreted by Engine. The node keeps only its configuration, loop-variable initialization, and static variable mapping. """ @@ -54,7 +54,7 @@ def version(cls) -> str: @override def _run( self, - ) -> Generator[NodeEventBase | LoopFrameRequest, None, None]: + ) -> Generator[NodeEventPayload | LoopFrameRequest, None, None]: loop_count = self.node_data.loop_count inputs: dict[str, object] = {"loop_count": loop_count} root_node_id = self.node_data.start_node_id @@ -84,7 +84,7 @@ def _resume_container_events( self, *, result: ContainerRunResult, - ) -> Generator[NodeEventBase | LoopFrameRequest, None, None]: + ) -> Generator[NodeEventPayload | LoopFrameRequest, None, None]: if isinstance(result, LoopFrameRequest): yield LoopNextEvent( index=result.index, diff --git a/src/graphon/nodes/parameter_extractor/parameter_extractor_node.py b/src/graphon/nodes/parameter_extractor/parameter_extractor_node.py index ab8e493c..51f97127 100644 --- a/src/graphon/nodes/parameter_extractor/parameter_extractor_node.py +++ b/src/graphon/nodes/parameter_extractor/parameter_extractor_node.py @@ -8,7 +8,7 @@ from dataclasses import dataclass from typing import Any, assert_never, overload, override -from graphon.entities.graph_init_params import GraphInitParams +from graphon.entities.graph_init_params import InitParams from graphon.enums import ( BuiltinNodeTypes, WorkflowNodeExecutionMetadataKey, @@ -44,7 +44,7 @@ LLMProtocol, PromptMessageSerializerProtocol, ) -from graphon.runtime.graph_runtime_state import GraphRuntimeState +from graphon.runtime.graph_runtime_state import RuntimeState from graphon.runtime.variable_pool import VariablePool from graphon.variables.factory import build_segment_with_type from graphon.variables.template_resolution import convert_template @@ -145,8 +145,8 @@ def __init__( node_id: str, data: ParameterExtractorNodeData, *, - graph_init_params: GraphInitParams, - graph_runtime_state: GraphRuntimeState, + graph_init_params: InitParams, + graph_runtime_state: RuntimeState, dependencies: _ParameterExtractorNodeDependencies, credentials_provider: object | None = None, model_factory: object | None = None, @@ -161,8 +161,8 @@ def __init__( node_id: str, data: ParameterExtractorNodeData, *, - graph_init_params: GraphInitParams, - graph_runtime_state: GraphRuntimeState, + graph_init_params: InitParams, + graph_runtime_state: RuntimeState, dependencies: None = None, credentials_provider: object | None = None, model_factory: object | None = None, @@ -177,8 +177,8 @@ def __init__( node_id: str, data: ParameterExtractorNodeData, *, - graph_init_params: GraphInitParams, - graph_runtime_state: GraphRuntimeState, + graph_init_params: InitParams, + graph_runtime_state: RuntimeState, dependencies: _ParameterExtractorNodeDependencies | None = None, credentials_provider: object | None = None, model_factory: object | None = None, diff --git a/src/graphon/nodes/question_classifier/question_classifier_node.py b/src/graphon/nodes/question_classifier/question_classifier_node.py index 70b44edb..63b7287a 100644 --- a/src/graphon/nodes/question_classifier/question_classifier_node.py +++ b/src/graphon/nodes/question_classifier/question_classifier_node.py @@ -6,7 +6,7 @@ from dataclasses import dataclass from typing import Any, override -from graphon.entities.graph_init_params import GraphInitParams +from graphon.entities.graph_init_params import InitParams from graphon.enums import ( BuiltinNodeTypes, NodeExecutionType, @@ -42,7 +42,7 @@ LLMProtocol, PromptMessageSerializerProtocol, ) -from graphon.runtime.graph_runtime_state import GraphRuntimeState +from graphon.runtime.graph_runtime_state import RuntimeState from graphon.template_rendering import Jinja2TemplateRenderer from graphon.utils.json_in_md_parser import parse_and_check_json_markdown from graphon.variables.template_resolution import convert_template @@ -108,8 +108,8 @@ def __init__( node_id: str, data: QuestionClassifierNodeData, *, - graph_init_params: GraphInitParams, - graph_runtime_state: GraphRuntimeState, + graph_init_params: InitParams, + graph_runtime_state: RuntimeState, dependencies: QuestionClassifierNodeDependencies | None = None, credentials_provider: object | None = None, model_factory: object | None = None, diff --git a/src/graphon/nodes/template_transform/template_transform_node.py b/src/graphon/nodes/template_transform/template_transform_node.py index 0f9eb395..b5b70e7f 100644 --- a/src/graphon/nodes/template_transform/template_transform_node.py +++ b/src/graphon/nodes/template_transform/template_transform_node.py @@ -5,13 +5,13 @@ from typing_extensions import TypeIs -from graphon.entities.graph_init_params import GraphInitParams +from graphon.entities.graph_init_params import InitParams from graphon.enums import BuiltinNodeTypes, WorkflowNodeExecutionStatus from graphon.node_events.base import NodeRunResult from graphon.nodes.base.entities import VariableSelector from graphon.nodes.base.node import Node from graphon.nodes.template_transform.entities import TemplateTransformNodeData -from graphon.runtime.graph_runtime_state import GraphRuntimeState +from graphon.runtime.graph_runtime_state import RuntimeState from graphon.template_rendering import ( Jinja2TemplateRenderer, TemplateRenderError, @@ -39,8 +39,8 @@ def __init__( node_id: str, data: TemplateTransformNodeData, *, - graph_init_params: GraphInitParams, - graph_runtime_state: GraphRuntimeState, + graph_init_params: InitParams, + graph_runtime_state: RuntimeState, jinja2_template_renderer: Jinja2TemplateRenderer, max_output_length: int | None = None, ) -> None: diff --git a/src/graphon/nodes/tool/tool_node.py b/src/graphon/nodes/tool/tool_node.py index 7c7b9ec4..796b54b8 100644 --- a/src/graphon/nodes/tool/tool_node.py +++ b/src/graphon/nodes/tool/tool_node.py @@ -4,7 +4,8 @@ from typing_extensions import TypeIs -from graphon.entities.graph_init_params import GraphInitParams +from graphon.engine_events.node import NodeRunStartedEvent +from graphon.entities.graph_init_params import InitParams from graphon.enums import ( BuiltinNodeTypes, WorkflowNodeExecutionMetadataKey, @@ -13,9 +14,8 @@ from graphon.file.enums import FileTransferMethod from graphon.file.file_factory import get_file_type_by_mime_type from graphon.file.models import File -from graphon.graph_events.node import NodeRunStartedEvent from graphon.node_events.base import ( - NodeEventBase, + NodeEventPayload, NodeRunResult, ) from graphon.node_events.node import ( @@ -31,7 +31,7 @@ ToolRuntimeMessage, ToolRuntimeParameter, ) -from graphon.runtime.graph_runtime_state import GraphRuntimeState +from graphon.runtime.graph_runtime_state import RuntimeState from graphon.runtime.variable_pool import VariablePool from graphon.variables.segments import ArrayFileSegment from graphon.variables.template_resolution import convert_template @@ -79,8 +79,8 @@ def __init__( node_id: str, data: ToolNodeData, *, - graph_init_params: GraphInitParams, - graph_runtime_state: GraphRuntimeState, + graph_init_params: InitParams, + graph_runtime_state: RuntimeState, tool_file_manager: ToolFileManagerProtocol, # TODO @-LAN: See https://github.com/langgenius/graphon/issues/new/choose. # ruff:ignore[line-contains-todo] # Make `runtime` optional once Graphon provides a default tool runtime @@ -107,7 +107,7 @@ def populate_start_event(self, event: NodeRunStartedEvent) -> None: event.provider_type = self.node_data.provider_type @override - def _run(self) -> Generator[NodeEventBase, None, None]: + def _run(self) -> Generator[NodeEventPayload, None, None]: """Run the tool node""" # fetch tool icon tool_info = { @@ -277,7 +277,7 @@ def _transform_message( node_id: str, tool_runtime: ToolRuntimeHandle, **_: Any, - ) -> Generator[NodeEventBase, None, None]: + ) -> Generator[NodeEventPayload, None, None]: """Convert graph-owned tool runtime messages into node outputs.""" state = _ToolMessageState() @@ -316,7 +316,7 @@ def transform_message( node_id: str, tool_runtime: ToolRuntimeHandle, **kwargs: Any, - ) -> Generator[NodeEventBase, None, None]: + ) -> Generator[NodeEventPayload, None, None]: """Convert tool runtime messages using the node's public test seam.""" yield from self._transform_message( messages=messages, @@ -333,7 +333,7 @@ def _dispatch_message( message: ToolRuntimeMessage, state: _ToolMessageState, node_id: str, - ) -> Generator[NodeEventBase, None, None]: + ) -> Generator[NodeEventPayload, None, None]: match message.type: case ( ToolRuntimeMessage.MessageType.IMAGE_LINK @@ -469,7 +469,7 @@ def _handle_linked_file_message( meta: Mapping[str, Any] | None, state: _ToolMessageState, **_: Any, - ) -> Generator[NodeEventBase, None, None]: + ) -> Generator[NodeEventPayload, None, None]: url = payload.text transfer_method = FileTransferMethod.TOOL_FILE tool_file_id: str | None = None @@ -507,7 +507,7 @@ def _handle_blob_message( meta: Mapping[str, Any] | None, state: _ToolMessageState, **_: Any, - ) -> Generator[NodeEventBase, None, None]: + ) -> Generator[NodeEventPayload, None, None]: tool_file_id = (meta or {}).get("tool_file_id") if isinstance(tool_file_id, str) and tool_file_id: self._resolve_tool_file( @@ -562,7 +562,7 @@ def _handle_blob_chunk_message( meta: Mapping[str, Any] | None, state: _ToolMessageState, **_: Any, - ) -> Generator[NodeEventBase, None, None]: + ) -> Generator[NodeEventPayload, None, None]: if not payload.id: msg = "tool blob chunk message is missing id" raise ToolFileError(msg) @@ -603,7 +603,7 @@ def _handle_text_message( state: _ToolMessageState, node_id: str, **_: Any, - ) -> Generator[NodeEventBase, None, None]: + ) -> Generator[NodeEventPayload, None, None]: state.text += payload.text yield StreamChunkEvent( selector=[node_id, "text"], @@ -617,7 +617,7 @@ def _handle_json_message( payload: ToolRuntimeMessage.JsonMessage, state: _ToolMessageState, **_: Any, - ) -> Generator[NodeEventBase, None, None]: + ) -> Generator[NodeEventPayload, None, None]: if payload.json_object: state.json_values.append(payload.json_object) yield from () @@ -630,7 +630,7 @@ def _handle_link_message( state: _ToolMessageState, node_id: str, **_: Any, - ) -> Generator[NodeEventBase, None, None]: + ) -> Generator[NodeEventPayload, None, None]: file_obj = (meta or {}).get("file") if isinstance(file_obj, File): state.files.append(file_obj) @@ -652,7 +652,7 @@ def _handle_variable_message( state: _ToolMessageState, node_id: str, **_: Any, - ) -> Generator[NodeEventBase, None, None]: + ) -> Generator[NodeEventPayload, None, None]: variable_name = payload.variable_name variable_value = payload.variable_value @@ -680,7 +680,7 @@ def _handle_file_message( meta: Mapping[str, Any], state: _ToolMessageState, **_: Any, - ) -> Generator[NodeEventBase, None, None]: + ) -> Generator[NodeEventPayload, None, None]: del payload if "file" not in meta: msg = "File message is missing 'file' key in meta" @@ -699,14 +699,14 @@ def _handle_log_message( *, payload: ToolRuntimeMessage.LogMessage, **_: Any, - ) -> Generator[NodeEventBase, None, None]: + ) -> Generator[NodeEventPayload, None, None]: del payload yield from () def _emit_final_stream_events( self, state: _ToolMessageState, - ) -> Generator[NodeEventBase, None, None]: + ) -> Generator[NodeEventPayload, None, None]: if state.blob_chunks: pending_ids = ", ".join(sorted(state.blob_chunks)) msg = f"tool blob chunk stream ended before completion: {pending_ids}" diff --git a/src/graphon/nodes/variable_assigner/v1/node.py b/src/graphon/nodes/variable_assigner/v1/node.py index 596a465c..84b68993 100644 --- a/src/graphon/nodes/variable_assigner/v1/node.py +++ b/src/graphon/nodes/variable_assigner/v1/node.py @@ -3,10 +3,10 @@ from collections.abc import Generator, Mapping, Sequence from typing import Any, assert_never, override -from graphon.entities.graph_init_params import GraphInitParams +from graphon.entities.graph_init_params import InitParams from graphon.enums import BuiltinNodeTypes, WorkflowNodeExecutionStatus from graphon.node_events.base import ( - NodeEventBase, + NodeEventPayload, NodeRunResult, ) from graphon.node_events.node import ( @@ -16,7 +16,7 @@ from graphon.nodes.base.node import Node from graphon.nodes.variable_assigner.common import helpers as common_helpers from graphon.nodes.variable_assigner.common.exc import VariableOperatorNodeError -from graphon.runtime.graph_runtime_state import GraphRuntimeState +from graphon.runtime.graph_runtime_state import RuntimeState from graphon.variables.types import SegmentType from graphon.variables.variables import ( ArrayAnyVariable, @@ -39,8 +39,8 @@ def __init__( node_id: str, data: VariableAssignerData, *, - graph_init_params: GraphInitParams, - graph_runtime_state: GraphRuntimeState, + graph_init_params: InitParams, + graph_runtime_state: RuntimeState, ) -> None: super().__init__( node_id=node_id, @@ -88,7 +88,7 @@ def _extract_variable_selector_to_variable_mapping( return mapping @override - def _run(self) -> Generator[NodeEventBase, None, None]: + def _run(self) -> Generator[NodeEventPayload, None, None]: assigned_variable_selector = self.node_data.assigned_variable_selector # Should be String, Number, Object, ArrayString, ArrayNumber, ArrayObject original_variable = self.graph_runtime_state.variable_pool.get_variable( diff --git a/src/graphon/nodes/variable_assigner/v2/node.py b/src/graphon/nodes/variable_assigner/v2/node.py index c1dfcddd..4d48e12d 100644 --- a/src/graphon/nodes/variable_assigner/v2/node.py +++ b/src/graphon/nodes/variable_assigner/v2/node.py @@ -4,10 +4,10 @@ from collections.abc import Generator, Mapping, MutableMapping, Sequence from typing import Any, override -from graphon.entities.graph_init_params import GraphInitParams +from graphon.entities.graph_init_params import InitParams from graphon.enums import BuiltinNodeTypes, WorkflowNodeExecutionStatus from graphon.node_events.base import ( - NodeEventBase, + NodeEventPayload, NodeRunResult, ) from graphon.node_events.node import ( @@ -17,7 +17,7 @@ from graphon.nodes.base.node import Node from graphon.nodes.variable_assigner.common import helpers as common_helpers from graphon.nodes.variable_assigner.common.exc import VariableOperatorNodeError -from graphon.runtime.graph_runtime_state import GraphRuntimeState +from graphon.runtime.graph_runtime_state import RuntimeState from graphon.variables.consts import SELECTORS_LENGTH from graphon.variables.types import SegmentType from graphon.variables.variables import VariableBase @@ -81,8 +81,8 @@ def __init__( node_id: str, data: VariableAssignerNodeData, *, - graph_init_params: GraphInitParams, - graph_runtime_state: GraphRuntimeState, + graph_init_params: InitParams, + graph_runtime_state: RuntimeState, ) -> None: super().__init__( node_id=node_id, @@ -134,7 +134,7 @@ def _extract_variable_selector_to_variable_mapping( return var_mapping @override - def _run(self) -> Generator[NodeEventBase, None, None]: + def _run(self) -> Generator[NodeEventPayload, None, None]: inputs = self.node_data.model_dump() process_data: dict[str, Any] = {} # NOTE: This node has no outputs diff --git a/src/graphon/runtime/__init__.py b/src/graphon/runtime/__init__.py index 4db967d5..0232912b 100644 --- a/src/graphon/runtime/__init__.py +++ b/src/graphon/runtime/__init__.py @@ -1,5 +1,5 @@ from .graph_runtime_state import ( - GraphRuntimeState, + RuntimeState, ) from .graph_runtime_state_protocol import ( ReadOnlyGraphRuntimeState, @@ -12,11 +12,11 @@ from .variable_pool import VariablePool, VariableValue __all__ = [ - "GraphRuntimeState", "ReadOnlyGraphRuntimeState", "ReadOnlyGraphRuntimeStateWrapper", "ReadOnlyVariablePool", "ReadOnlyVariablePoolWrapper", + "RuntimeState", "VariablePool", "VariableValue", ] diff --git a/src/graphon/runtime/container_state.py b/src/graphon/runtime/container_state.py index d9aec4a3..d9b015ff 100644 --- a/src/graphon/runtime/container_state.py +++ b/src/graphon/runtime/container_state.py @@ -11,6 +11,7 @@ from graphon.nodes.container_effects import ( ContainerAwaitRequest, ContainerValue, + CustomContainerRequest, IterationFrameRequest, LoopFrameRequest, ) @@ -81,8 +82,21 @@ class IterationRunState(BaseModel): errors: tuple[str, ...] = () +class CustomContainerRunState(BaseModel): + """Serializable state for one custom container invocation.""" + + model_config = ConfigDict(frozen=True) + + kind: Literal["custom"] = "custom" + invocation_id: str + frame_id: str + node_id: str + started_at: datetime + payload: str + + ContainerRunState = Annotated[ - LoopRunState | IterationRunState, + LoopRunState | IterationRunState | CustomContainerRunState, Field(discriminator="kind"), ] @@ -122,6 +136,14 @@ def create_container_run_state( flatten_output=request.flatten_output, parallel_nums=request.parallel_nums, ) + case CustomContainerRequest(): + return CustomContainerRunState( + invocation_id=invocation_id, + frame_id=frame_id, + node_id=node_id, + started_at=started_at, + payload=request.payload, + ) case _: assert_never(request) @@ -157,7 +179,18 @@ class IterationFrameState(BaseModel): runtime_data: FrameRuntimeData +class CustomContainerFrameState(BaseModel): + """Serializable state for one custom container child frame.""" + + model_config = ConfigDict(frozen=True) + + kind: Literal["custom"] = "custom" + frame_id: str + parent_invocation_id: str + runtime_data: FrameRuntimeData + + ContainerFrameState = Annotated[ - LoopFrameState | IterationFrameState, + LoopFrameState | IterationFrameState | CustomContainerFrameState, Field(discriminator="kind"), ] diff --git a/src/graphon/graph_engine/domain/graph_execution.py b/src/graphon/runtime/execution.py similarity index 79% rename from src/graphon/graph_engine/domain/graph_execution.py rename to src/graphon/runtime/execution.py index bef86961..cc822cc9 100644 --- a/src/graphon/graph_engine/domain/graph_execution.py +++ b/src/graphon/runtime/execution.py @@ -1,17 +1,28 @@ -"""GraphExecution aggregate root managing the overall graph execution state.""" +"""Graph-wide execution state shared by every runtime frame.""" from __future__ import annotations from dataclasses import dataclass, field -from typing import Annotated, Literal +from typing import Annotated, Final, Literal from uuid import uuid4 from pydantic import BaseModel, Field, TypeAdapter from graphon.entities.pause_reason import PauseReason -from graphon.graph_engine.ready_queue.protocol import ROOT_FRAME_ID -from .node_execution import NodeExecution +ROOT_FRAME_ID: Final = "root" + + +@dataclass +class NodeExecution: + """Mutable retry and identity state for one frame-local node.""" + + execution_id: str + retry_count: int = 0 + + def increment_retry(self) -> None: + """Increment the retry count for this node.""" + self.retry_count += 1 class GraphExecutionErrorStateV1(BaseModel): @@ -173,8 +184,22 @@ def dumps(self) -> str: return state.model_dump_json() - def loads(self, data: str) -> None: - """Restore aggregate state from a serialized JSON string.""" + @classmethod + def from_snapshot(cls, data: str) -> GraphExecution: + """Construct a complete execution state from serialized JSON. + + Both current version 2 snapshots and legacy version 1 snapshots are + accepted. Legacy node executions are assigned to the root frame, and a + missing legacy execution ID is regenerated exactly as before. The JSON + field names and version values remain unchanged. + + Args: + data: Serialized graph execution snapshot. + + Returns: + A fully initialized execution state ready for runtime use. + + """ serialized_state = _GRAPH_EXECUTION_STATE_ADAPTER.validate_json(data) if isinstance(serialized_state, GraphExecutionStateV1): state = GraphExecutionState( @@ -208,24 +233,23 @@ def loads(self, data: str) -> None: else: state = serialized_state - if self.workflow_id != state.workflow_id: - msg = "Serialized workflow_id does not match aggregate identity" - raise ValueError(msg) - - self.started = state.started - self.completed = state.completed - self.aborted = state.aborted - self.paused = state.paused - self.pause_reasons = state.pause_reasons - self.error = RuntimeError(state.error) if state.error is not None else None - self.exceptions_count = state.exceptions_count - self.node_executions = { - (item.frame_id, item.node_id): NodeExecution( - retry_count=item.retry_count, - execution_id=item.execution_id, - ) - for item in state.node_executions - } + return cls( + workflow_id=state.workflow_id, + started=state.started, + completed=state.completed, + aborted=state.aborted, + paused=state.paused, + pause_reasons=state.pause_reasons, + error=RuntimeError(state.error) if state.error is not None else None, + exceptions_count=state.exceptions_count, + node_executions={ + (item.frame_id, item.node_id): NodeExecution( + retry_count=item.retry_count, + execution_id=item.execution_id, + ) + for item in state.node_executions + }, + ) def record_node_failure(self) -> None: """Increment the count of node failures encountered during execution.""" diff --git a/src/graphon/runtime/graph_runtime_state.py b/src/graphon/runtime/graph_runtime_state.py index 6ad16394..ebaf463a 100644 --- a/src/graphon/runtime/graph_runtime_state.py +++ b/src/graphon/runtime/graph_runtime_state.py @@ -1,6 +1,5 @@ from __future__ import annotations -import json import threading from abc import abstractmethod from collections.abc import Callable, Mapping, Sequence @@ -20,89 +19,10 @@ from graphon.runtime.ready_queue import ReadyQueue from graphon.runtime.variable_pool import VariablePool -if TYPE_CHECKING: - from graphon.entities.pause_reason import PauseReason - from graphon.graph_engine.ready_queue import ReadyTask - - -class NodeExecutionProtocol(Protocol): - """Structural interface for persisted per-node execution state.""" - - retry_count: int - execution_id: str - - @abstractmethod - def increment_retry(self) -> None: - """Increment the retry counter for the node execution.""" - ... - - -class GraphExecutionProtocol(Protocol): - """Structural interface for graph execution aggregate. - - Defines the minimal set of attributes and methods required - from a GraphExecution entity for runtime orchestration and - state management. - """ - - workflow_id: str - started: bool - completed: bool - aborted: bool - paused: bool - error: Exception | None - exceptions_count: int - pause_reasons: list[PauseReason] - - @abstractmethod - def start(self) -> None: - """Transition execution into the running state.""" - ... - - @abstractmethod - def complete(self) -> None: - """Mark execution as successfully completed.""" - ... - - @abstractmethod - def abort(self, reason: str) -> None: - """Abort execution in response to an external stop request.""" - ... - - @abstractmethod - def pause(self, reason: PauseReason) -> None: - """Pause execution with a recorded reason.""" - ... - - @abstractmethod - def fail(self, error: Exception) -> None: - """Record an unrecoverable error and end execution.""" - ... - - @abstractmethod - def record_node_failure(self) -> None: - """Increment the count of node failures observed during execution.""" - ... - - @abstractmethod - def get_or_create_node_execution( - self, - *, - frame_id: str, - node_id: str, - ) -> NodeExecutionProtocol: - """Return the execution entity for a task, creating it when needed.""" - ... +from .execution import ROOT_FRAME_ID, GraphExecution - @abstractmethod - def dumps(self) -> str: - """Serialize execution state into a JSON payload.""" - ... - - @abstractmethod - def loads(self, data: str) -> None: - """Restore execution state from a previously serialized payload.""" - ... +if TYPE_CHECKING: + from graphon.engine.ready_queue import ReadyTask class NodeProtocol(Protocol): @@ -181,7 +101,7 @@ class _GraphRuntimeStateSnapshot(BaseModel): model_config = ConfigDict(frozen=True) - version: Literal["2.0"] + version: Literal["2.0", "3.0"] start_at: float node_run_steps: int = Field(ge=0) llm_usage: LLMUsage @@ -205,29 +125,21 @@ class _GraphRuntimeStateSnapshot(BaseModel): def _new_ready_queue() -> ReadyQueue: - from graphon.graph_engine.ready_queue import ( # ruff:ignore[import-outside-top-level] + from graphon.engine.ready_queue import ( # ruff:ignore[import-outside-top-level] InMemoryReadyQueue, ) return InMemoryReadyQueue() -def _new_graph_execution(workflow_id: str = "") -> GraphExecutionProtocol: - from graphon.graph_engine.domain.graph_execution import ( # ruff:ignore[import-outside-top-level] - GraphExecution, - ) - - return GraphExecution(workflow_id=workflow_id) - - -class GraphRuntimeState: # ruff:ignore[too-many-public-methods] +class RuntimeState: # ruff:ignore[too-many-public-methods] """Mutable runtime state shared across graph execution components. - `GraphRuntimeState` encapsulates the runtime state of workflow execution, + `RuntimeState` encapsulates the runtime state of workflow execution, including scheduling details, variable values, and timing information. Values that are initialized prior to workflow execution and remain constant - throughout the execution should be part of `GraphInitParams` instead. + throughout the execution should be part of `InitParams` instead. """ _container_state_lock: threading.Lock @@ -242,12 +154,46 @@ def __init__( node_run_steps: int = 0, ready_queue: ReadyQueue | None = None, deferred_ready_queue: ReadyQueue | None = None, - graph_execution: GraphExecutionProtocol | None = None, + workflow_id: str | None = None, + graph_execution: GraphExecution | None = None, execution_context: AbstractContextManager[object] | None = None, ) -> None: + """Initialize all mutable state owned by one graph execution frame. + + A newly started root runtime supplies ``workflow_id`` so this constructor + can create the execution aggregate. Child frames and restored runtimes + instead supply their existing ``graph_execution`` so every frame shares + one workflow identity and lifecycle. Callers may supply both forms when + useful at an API boundary, but their workflow IDs must agree. + + Args: + variable_pool: Variables visible to nodes in this frame. + start_at: Unix timestamp at which execution started. + llm_usage: Accumulated language-model usage, copied on input. + outputs: Current workflow outputs, copied on input. + node_run_steps: Number of node runs already completed. + ready_queue: Queue for runnable tasks, or a local queue by default. + deferred_ready_queue: Queue for tasks held while execution is paused. + workflow_id: Identity used to create a new execution aggregate. + graph_execution: Existing aggregate shared by child or restored frames. + execution_context: Context entered by workers around node execution. + + Raises: + ValueError: If ``node_run_steps`` is negative, neither identity form + is supplied, or the two supplied workflow identities disagree. + + """ if node_run_steps < 0: msg = "node_run_steps must be non-negative" raise ValueError(msg) + if graph_execution is None: + if workflow_id is None: + msg = "workflow_id or graph_execution is required" + raise ValueError(msg) + graph_execution = GraphExecution(workflow_id=workflow_id) + elif workflow_id is not None and workflow_id != graph_execution.workflow_id: + msg = "workflow_id must match graph_execution.workflow_id" + raise ValueError(msg) self._variable_pool = variable_pool self._start_at = start_at self._llm_usage = ( @@ -264,9 +210,7 @@ def __init__( if deferred_ready_queue is not None else _new_ready_queue() ) - self._graph_execution = ( - graph_execution if graph_execution is not None else _new_graph_execution() - ) + self._graph_execution = graph_execution self._execution_context = ( execution_context if execution_context is not None else nullcontext() ) @@ -274,6 +218,7 @@ def __init__( self._container_frames: dict[str, ContainerFrameState] = {} self._pending_graph_node_states: dict[str, NodeState] = {} self._pending_graph_edge_states: dict[str, NodeState] = {} + self._has_pending_graph_state = False self._container_state_lock = threading.Lock() @property @@ -289,7 +234,7 @@ def deferred_ready_queue(self) -> ReadyQueue: return self._deferred_ready_queue @property - def graph_execution(self) -> GraphExecutionProtocol: + def graph_execution(self) -> GraphExecution: return self._graph_execution @property @@ -344,11 +289,46 @@ def node_run_steps(self) -> int: def increment_node_run_steps(self) -> None: self._node_run_steps += 1 + def restore_graph_state( + self, + *, + node_states: Mapping[str, NodeState], + edge_states: Mapping[str, NodeState], + ) -> None: + """Stage persisted graph states for the graph attached to this runtime. + + Frame restoration constructs runtime state before constructing its + scoped graph. This method records the complete node and edge mappings; + :meth:`attach_graph` validates their topology and applies them atomically + before the frame becomes executable. Staging is rejected after graph + attachment so callers cannot overwrite a live graph accidentally. + + Args: + node_states: Persisted node states keyed by node ID. + edge_states: Persisted edge states keyed by edge ID. + + Raises: + RuntimeError: If a graph is already attached to this runtime state. + + """ + if self._graph is not None or self._has_pending_graph_state: + msg = "graph state must be restored before attaching a graph" + raise RuntimeError(msg) + self._pending_graph_node_states = dict(node_states) + self._pending_graph_edge_states = dict(edge_states) + self._has_pending_graph_state = True + def attach_graph(self, graph: GraphProtocol) -> None: """Attach the materialized graph to the runtime state.""" if self._graph is not None and self._graph is not graph: - msg = "GraphRuntimeState already attached to a different graph instance" + msg = "RuntimeState already attached to a different graph instance" raise ValueError(msg) + if self._has_pending_graph_state and ( + set(self._pending_graph_node_states) != set(graph.nodes) + or set(self._pending_graph_edge_states) != set(graph.edges) + ): + msg = "Saved graph state does not match rebuilt graph" + raise RuntimeError(msg) self._graph = graph self._apply_pending_graph_state() @@ -361,6 +341,7 @@ def _apply_pending_graph_state(self) -> None: self._graph.edges[edge_id].state = state self._pending_graph_node_states.clear() self._pending_graph_edge_states.clear() + self._has_pending_graph_state = False def dumps(self) -> str: """Serialize runtime state into a JSON string.""" @@ -378,7 +359,7 @@ def dumps(self) -> str: edge_id: edge.state for edge_id, edge in self._graph.edges.items() } return _GraphRuntimeStateSnapshot( - version="2.0", + version="3.0", start_at=self._start_at, node_run_steps=self._node_run_steps, llm_usage=self._llm_usage, @@ -395,11 +376,11 @@ def dumps(self) -> str: @classmethod def from_snapshot( - cls: type[GraphRuntimeState], + cls: type[RuntimeState], data: str, *, ready_queue_factory: Callable[[], ReadyQueue] = _new_ready_queue, - ) -> GraphRuntimeState: + ) -> RuntimeState: """Restore runtime state from a serialized snapshot.""" snapshot = _GRAPH_RUNTIME_STATE_SNAPSHOT_ADAPTER.validate_json(data) @@ -407,8 +388,7 @@ def from_snapshot( ready_queue.loads(snapshot.ready_queue) deferred_ready_queue = ready_queue_factory() if isinstance(snapshot, _GraphRuntimeStateSnapshotV1): - from graphon.graph_engine.ready_queue import ( # ruff:ignore[import-outside-top-level] - ROOT_FRAME_ID, + from graphon.engine.ready_queue import ( # ruff:ignore[import-outside-top-level] StartTask, ) @@ -430,11 +410,7 @@ def from_snapshot( graph_node_states = snapshot.graph_node_states graph_edge_states = snapshot.graph_edge_states - execution_payload = json.loads(snapshot.graph_execution) - graph_execution = _new_graph_execution( - workflow_id=execution_payload["workflow_id"], - ) - graph_execution.loads(snapshot.graph_execution) + graph_execution = GraphExecution.from_snapshot(snapshot.graph_execution) state = cls( variable_pool=snapshot.variable_pool, @@ -448,8 +424,10 @@ def from_snapshot( ) state._container_runs = {run.invocation_id: run for run in container_runs} state._container_frames = {frame.frame_id: frame for frame in container_frames} - state._pending_graph_node_states = graph_node_states - state._pending_graph_edge_states = graph_edge_states + state.restore_graph_state( + node_states=graph_node_states, + edge_states=graph_edge_states, + ) return state def defer_ready_task(self, task: ReadyTask) -> None: diff --git a/src/graphon/runtime/graph_runtime_state_protocol.py b/src/graphon/runtime/graph_runtime_state_protocol.py index f4e1e31e..6ef13f95 100644 --- a/src/graphon/runtime/graph_runtime_state_protocol.py +++ b/src/graphon/runtime/graph_runtime_state_protocol.py @@ -21,7 +21,7 @@ def get_by_prefix(self, prefix: str, /) -> Mapping[str, object]: class ReadOnlyGraphRuntimeState(Protocol): - """Read-only view of GraphRuntimeState for layers. + """Read-only view of RuntimeState for layers. This protocol defines a read-only interface that prevents layers from modifying the graph runtime state while still allowing observation. diff --git a/src/graphon/runtime/read_only_wrappers.py b/src/graphon/runtime/read_only_wrappers.py index 822c3ce5..8c2b55f5 100644 --- a/src/graphon/runtime/read_only_wrappers.py +++ b/src/graphon/runtime/read_only_wrappers.py @@ -6,7 +6,7 @@ from graphon.model_runtime.entities.llm_entities import LLMUsage from graphon.variables.segments import Segment -from .graph_runtime_state import GraphRuntimeState +from .graph_runtime_state import RuntimeState from .graph_runtime_state_protocol import ReadOnlyVariablePool from .variable_pool import VariablePool @@ -28,9 +28,9 @@ def get_by_prefix(self, prefix: str, /) -> Mapping[str, object]: class ReadOnlyGraphRuntimeStateWrapper: - """Expose a defensive, read-only view of ``GraphRuntimeState``.""" + """Expose a defensive, read-only view of ``RuntimeState``.""" - def __init__(self, state: GraphRuntimeState) -> None: + def __init__(self, state: RuntimeState) -> None: self._state = state self._variable_pool_wrapper = ReadOnlyVariablePoolWrapper(state.variable_pool) diff --git a/src/graphon/runtime/ready_queue.py b/src/graphon/runtime/ready_queue.py index d1c6c593..ac00a119 100644 --- a/src/graphon/runtime/ready_queue.py +++ b/src/graphon/runtime/ready_queue.py @@ -6,7 +6,7 @@ from typing import TYPE_CHECKING, Protocol if TYPE_CHECKING: - from graphon.graph_engine.ready_queue.protocol import ReadyTask + from graphon.engine.ready_queue.entities import ReadyTask class ReadyQueue(Protocol): diff --git a/tests/dsl/test_node_factory.py b/tests/dsl/test_node_factory.py index 5f921529..f86a06e1 100644 --- a/tests/dsl/test_node_factory.py +++ b/tests/dsl/test_node_factory.py @@ -17,14 +17,14 @@ PluginDependencyType, ) from graphon.dsl.errors import DslError -from graphon.entities.graph_config import NodeConfigDict -from graphon.file.enums import FileTransferMethod, FileType -from graphon.file.models import File -from graphon.graph_events.node import ( +from graphon.engine_events.node import ( NodeRunFailedEvent, NodeRunSucceededEvent, NodeRunVariableUpdatedEvent, ) +from graphon.entities.graph_config import NodeConfigDict +from graphon.file.enums import FileTransferMethod, FileType +from graphon.file.models import File from graphon.http import HttpResponse from graphon.model_runtime.entities.common_entities import I18nObject from graphon.model_runtime.entities.llm_entities import LLMResult, LLMUsage @@ -55,7 +55,7 @@ from graphon.nodes.variable_assigner.v2.node import ( VariableAssignerNode as VariableAssignerNodeV2, ) -from graphon.runtime.graph_runtime_state import GraphRuntimeState +from graphon.runtime.graph_runtime_state import RuntimeState from tests.helpers import build_graph_init_params, build_variable_pool _OPENAI_PLUGIN_ID = "langgenius/openai:0.3.8@test" @@ -364,7 +364,8 @@ def _dsl_node_factory( graph_init_params=build_graph_init_params( graph_config={"nodes": [], "edges": []}, ), - graph_runtime_state=GraphRuntimeState( + graph_runtime_state=RuntimeState( + workflow_id="workflow", variable_pool=build_variable_pool(variables=variables), start_at=0, ), @@ -474,11 +475,13 @@ def test_slim_dsl_node_factory_rebinds_graph_runtime_state() -> None: "edges": [], } graph_init_params = build_graph_init_params(graph_config=graph_config) - original_runtime_state = GraphRuntimeState( + original_runtime_state = RuntimeState( + workflow_id="workflow", variable_pool=build_variable_pool(variables=[(["start", "query"], "before")]), start_at=1, ) - rebound_runtime_state = GraphRuntimeState( + rebound_runtime_state = RuntimeState( + workflow_id="workflow", variable_pool=build_variable_pool(variables=[(["start", "query"], "after")]), start_at=2, ) diff --git a/src/graphon/graph_engine/entities/__init__.py b/tests/engine/__init__.py similarity index 100% rename from src/graphon/graph_engine/entities/__init__.py rename to tests/engine/__init__.py diff --git a/tests/graph_engine/test_cooperative_container_execution.py b/tests/engine/test_cooperative_container_execution.py similarity index 82% rename from tests/graph_engine/test_cooperative_container_execution.py rename to tests/engine/test_cooperative_container_execution.py index bc07a7ff..e54bdf04 100644 --- a/tests/graph_engine/test_cooperative_container_execution.py +++ b/tests/engine/test_cooperative_container_execution.py @@ -7,36 +7,33 @@ import pytest -from graphon.enums import ( - BuiltinNodeTypes, - NodeExecutionType, - WorkflowNodeExecutionStatus, +from graphon.engine.event.node_failure import NodeFailureHandler +from graphon.engine.frame import ExecutionFrame, FrameRegistry +from graphon.engine.layer import Layer +from graphon.engine.ready_queue.entities import ( + ResumeTask, + StartTask, ) -from graphon.graph.graph import Graph -from graphon.graph_engine.domain.graph_execution import GraphExecution -from graphon.graph_engine.entities.tasks import ( +from graphon.engine.ready_queue.in_memory import InMemoryReadyQueue +from graphon.engine.scheduler import Scheduler +from graphon.engine.worker import ( ContainerAwaitTask, DispatchTask, - TaskEvent, -) -from graphon.graph_engine.error_handler import ErrorHandler -from graphon.graph_engine.frames import ExecutionFrame, FrameRegistry -from graphon.graph_engine.graph_state_manager import GraphStateManager -from graphon.graph_engine.graph_traversal.edge_processor import EdgeProcessor -from graphon.graph_engine.graph_traversal.skip_propagator import SkipPropagator -from graphon.graph_engine.layers.base import GraphEngineLayer -from graphon.graph_engine.ready_queue.in_memory import InMemoryReadyQueue -from graphon.graph_engine.ready_queue.protocol import ( - ResumeTask, - StartTask, + NodeEventTask, + Worker, ) -from graphon.graph_engine.worker import Worker -from graphon.graph_events.base import GraphEngineEvent, GraphNodeEventBase -from graphon.graph_events.node import ( +from graphon.engine_events.base import EngineEvent, NodeEvent +from graphon.engine_events.node import ( NodeRunFailedEvent, NodeRunStartedEvent, NodeRunSucceededEvent, ) +from graphon.enums import ( + BuiltinNodeTypes, + NodeExecutionType, + WorkflowNodeExecutionStatus, +) +from graphon.graph.graph import Graph from graphon.node_events.base import NodeRunResult from graphon.nodes.base.node import Node from graphon.nodes.container_effects import ( @@ -47,7 +44,8 @@ build_container_value, ) from graphon.runtime.container_state import create_container_run_state -from graphon.runtime.graph_runtime_state import GraphRuntimeState +from graphon.runtime.execution import GraphExecution +from graphon.runtime.graph_runtime_state import RuntimeState from graphon.runtime.variable_pool import VariablePool @@ -55,24 +53,15 @@ def _execution_frame( *, frame_id: str, graph: Graph, - graph_runtime_state: GraphRuntimeState, + graph_runtime_state: RuntimeState, ) -> ExecutionFrame: - state_manager = GraphStateManager(graph, graph_runtime_state, frame_id) - skip_propagator = SkipPropagator( - graph=graph, - state_manager=state_manager, - ) + scheduler = Scheduler(graph, graph_runtime_state, frame_id) return ExecutionFrame( frame_id=frame_id, graph=graph, - graph_runtime_state=graph_runtime_state, - state_manager=state_manager, - edge_processor=EdgeProcessor( - graph=graph, - state_manager=state_manager, - skip_propagator=skip_propagator, - ), - error_handler=ErrorHandler(graph, graph_runtime_state.graph_execution), + state=graph_runtime_state, + scheduler=scheduler, + failure_handler=NodeFailureHandler(graph, graph_runtime_state.graph_execution), ) @@ -88,15 +77,15 @@ def _container_result() -> ContainerExecutionResult: ) -class _RecordingLayer(GraphEngineLayer): +class _RecordingLayer(Layer): def __init__(self) -> None: super().__init__() - self.end_events: list[GraphNodeEventBase | None] = [] + self.end_events: list[NodeEvent | None] = [] def on_graph_start(self) -> None: return - def on_event(self, event: GraphEngineEvent) -> None: + def on_event(self, event: EngineEvent) -> None: _ = event def on_graph_end(self, error: Exception | None) -> None: @@ -106,7 +95,7 @@ def on_node_run_end( self, node: object, error: Exception | None, - result_event: GraphNodeEventBase | None = None, + result_event: NodeEvent | None = None, ) -> None: _ = node _ = error @@ -220,7 +209,7 @@ def bind_execution_id(self, execution_id: str) -> None: def run( self, - ) -> Generator[GraphNodeEventBase | LoopFrameRequest, object, None]: + ) -> Generator[NodeEvent | LoopFrameRequest, object, None]: started_at = datetime.now(UTC).replace(tzinfo=None) yield NodeRunStartedEvent( id=self.execution_id, @@ -246,7 +235,7 @@ def resume_container( *, result: ContainerRunResult, started_at: datetime, - ) -> Generator[GraphNodeEventBase | LoopFrameRequest, None, None]: + ) -> Generator[NodeEvent | LoopFrameRequest, None, None]: assert isinstance(result, ContainerExecutionResult) node_run_result = NodeRunResult( status=result.node_run_result.status, @@ -275,9 +264,10 @@ def resume_container( ) ready_queue = InMemoryReadyQueue() ready_queue.put(StartTask(frame_id="root", node_id="loop")) - event_queue: queue.Queue[DispatchTask] = queue.Queue() + dispatch_queue: queue.Queue[DispatchTask] = queue.Queue() graph_execution = GraphExecution(workflow_id="workflow") - runtime_state = GraphRuntimeState( + runtime_state = RuntimeState( + workflow_id="workflow", variable_pool=VariablePool(), start_at=1, ready_queue=ready_queue, @@ -296,7 +286,7 @@ def resume_container( task_claiming.set() worker = Worker( ready_queue=ready_queue, - event_queue=event_queue, + dispatch_queue=dispatch_queue, frame_registry=frame_registry, layers=[layer], task_claim_lock=Lock(), @@ -305,10 +295,10 @@ def resume_container( worker.start() try: - started = event_queue.get(timeout=1) - assert isinstance(started, TaskEvent) + started = dispatch_queue.get(timeout=1) + assert isinstance(started, NodeEventTask) assert isinstance(started.event, NodeRunStartedEvent) - await_task = event_queue.get(timeout=1) + await_task = dispatch_queue.get(timeout=1) assert isinstance(await_task, ContainerAwaitTask) assert container_node.await_was_reached assert not container_node.body_after_await_was_consumed @@ -329,12 +319,12 @@ def resume_container( result=_container_result(), ), ) - succeeded = event_queue.get(timeout=1) + succeeded = dispatch_queue.get(timeout=1) finally: worker.stop() worker.join(timeout=1) - assert isinstance(succeeded, TaskEvent) + assert isinstance(succeeded, NodeEventTask) assert isinstance(succeeded.event, NodeRunSucceededEvent) assert succeeded.event.node_run_result.outputs == {"answer": "ok"} with pytest.raises(KeyError): @@ -357,7 +347,7 @@ def resume_container( *, result: ContainerRunResult, started_at: datetime, - ) -> Generator[GraphNodeEventBase | LoopFrameRequest, None, None]: + ) -> Generator[NodeEvent | LoopFrameRequest, None, None]: _ = result _ = started_at if False: @@ -366,9 +356,10 @@ def resume_container( raise RuntimeError(msg) ready_queue = InMemoryReadyQueue() - event_queue: queue.Queue[DispatchTask] = queue.Queue() + dispatch_queue: queue.Queue[DispatchTask] = queue.Queue() graph_execution = GraphExecution(workflow_id="workflow") - runtime_state = GraphRuntimeState( + runtime_state = RuntimeState( + workflow_id="workflow", variable_pool=VariablePool(), start_at=1, ready_queue=ready_queue, @@ -424,7 +415,7 @@ def resume_container( task_claiming.set() worker = Worker( ready_queue=ready_queue, - event_queue=event_queue, + dispatch_queue=dispatch_queue, frame_registry=frame_registry, layers=[], task_claim_lock=Lock(), @@ -433,12 +424,12 @@ def resume_container( worker.start() try: - failed = event_queue.get(timeout=1) + failed = dispatch_queue.get(timeout=1) finally: worker.stop() worker.join(timeout=1) - assert isinstance(failed, TaskEvent) + assert isinstance(failed, NodeEventTask) assert failed.frame_id == "parent-frame" assert isinstance(failed.event, NodeRunFailedEvent) assert failed.event.error == "resume bad" diff --git a/tests/graph_engine/test_dispatch_patterns.py b/tests/engine/test_dispatch_patterns.py similarity index 70% rename from tests/graph_engine/test_dispatch_patterns.py rename to tests/engine/test_dispatch_patterns.py index 03e0f4eb..79c93fb6 100644 --- a/tests/graph_engine/test_dispatch_patterns.py +++ b/tests/engine/test_dispatch_patterns.py @@ -6,58 +6,61 @@ from time import time from types import SimpleNamespace from typing import Any, ClassVar, cast -from unittest.mock import MagicMock +from unittest.mock import MagicMock, call import pytest -from graphon.entities.graph_init_params import GraphInitParams -from graphon.entities.pause_reason import HitlRequired -from graphon.enums import ( - BuiltinNodeTypes, - ErrorHandleMode, - NodeExecutionType, - NodeState, - NodeType, -) -from graphon.graph.graph import Graph -from graphon.graph_engine.command_channels.redis_channel import RedisChannel -from graphon.graph_engine.config import GraphEngineConfig -from graphon.graph_engine.container_handlers import ContainerHandler -from graphon.graph_engine.domain.graph_execution import GraphExecution -from graphon.graph_engine.entities.commands import ( +from graphon.engine import Engine +from graphon.engine.command import ( AbortCommand, - CommandType, + CommandProcessor, PauseCommand, UpdateVariablesCommand, ) -from graphon.graph_engine.entities.tasks import ( - DispatchTask, - TaskEvent, -) -from graphon.graph_engine.event_management.event_handlers import EventHandler -from graphon.graph_engine.event_management.event_manager import EventManager -from graphon.graph_engine.frames import ExecutionFrame, FrameRegistry -from graphon.graph_engine.graph_state_manager import GraphStateManager -from graphon.graph_engine.iteration_container_handler import IterationContainerHandler -from graphon.graph_engine.layers.execution_limits import ( - ExecutionLimitsLayer, - LimitType, +from graphon.engine.command.builtin.in_memory import InMemoryChannel +from graphon.engine.command.builtin.redis import RedisChannel +from graphon.engine.container_handler import ( + ContainerHandler, + IterationContainerHandler, + LoopContainerHandler, ) -from graphon.graph_engine.loop_container_handler import LoopContainerHandler -from graphon.graph_engine.orchestration.dispatcher import Dispatcher -from graphon.graph_engine.ready_queue.in_memory import InMemoryReadyQueue -from graphon.graph_engine.ready_queue.protocol import ( +from graphon.engine.dispatcher import Dispatcher +from graphon.engine.event.processor import NodeEventProcessor +from graphon.engine.event.stream import EventStream +from graphon.engine.frame import ExecutionFrame, FrameRegistry +from graphon.engine.layer import ExecutionLimitsLayer +from graphon.engine.ready_queue.entities import ( ReadyTask, ResumeTask, StartTask, ) -from graphon.graph_engine.worker import Worker -from graphon.graph_engine.worker_management import WorkerPool -from graphon.graph_events.node import ( +from graphon.engine.ready_queue.in_memory import InMemoryReadyQueue +from graphon.engine.scheduler import Scheduler +from graphon.engine.worker import ( + DispatchTask, + NodeEventTask, + Worker, + WorkerPool, +) +from graphon.engine_events.node import ( NodeRunPauseRequestedEvent, NodeRunStartedEvent, NodeRunSucceededEvent, ) +from graphon.engine_events.traversal import ( + GraphEdgeSkippedEvent, + GraphEdgeTakenEvent, +) +from graphon.entities.graph_init_params import InitParams +from graphon.entities.pause_reason import HitlRequired, SchedulingPause +from graphon.enums import ( + BuiltinNodeTypes, + ErrorHandleMode, + NodeExecutionType, + NodeState, + NodeType, +) +from graphon.graph.graph import Graph from graphon.model_runtime.entities.llm_entities import LLMUsage from graphon.node_events.base import NodeRunResult from graphon.nodes.container_effects import ( @@ -74,12 +77,14 @@ IterationRunState, create_container_run_state, ) -from graphon.runtime.graph_runtime_state import GraphRuntimeState +from graphon.runtime.execution import GraphExecution +from graphon.runtime.graph_runtime_state import RuntimeState from graphon.runtime.variable_pool import VariablePool +from graphon.variables.variables import StringVariable def _variable_value( - runtime_state: GraphRuntimeState, + runtime_state: RuntimeState, selector: list[str], ) -> object: variable = runtime_state.variable_pool.get(selector) @@ -92,58 +97,58 @@ def _execution_frame( frame_id: str, graph: Graph, graph_runtime_state: object | None = None, - state_manager: object | None = None, - edge_processor: object | None = None, - error_handler: object | None = None, + scheduler: object | None = None, + failure_handler: object | None = None, + container_id: str = "", ) -> ExecutionFrame: if isinstance(graph_runtime_state, MagicMock): graph_runtime_state.has_container_frame.return_value = False - if edge_processor is None: - resolved_edge_processor = MagicMock() - resolved_edge_processor.process_node_success.return_value = ([], []) - resolved_edge_processor.handle_branch_completion.return_value = ([], []) + if scheduler is None: + resolved_scheduler = MagicMock() + resolved_scheduler.process_node_success.return_value = ([], []) + resolved_scheduler.handle_branch_completion.return_value = ([], []) else: - resolved_edge_processor = edge_processor + resolved_scheduler = scheduler return ExecutionFrame( frame_id=frame_id, graph=graph, - graph_runtime_state=cast(Any, graph_runtime_state or MagicMock()), - state_manager=cast(Any, state_manager or MagicMock()), - edge_processor=cast(Any, resolved_edge_processor), - error_handler=cast(Any, error_handler or MagicMock()), + state=cast(Any, graph_runtime_state or MagicMock()), + scheduler=cast(Any, resolved_scheduler), + failure_handler=cast(Any, failure_handler or MagicMock()), + container_id=container_id, ) -def _event_handler( +def _event_processor( *, graph_execution: object, - event_collector: object, + event_stream: object, frame_registry: FrameRegistry, -) -> EventHandler: +) -> NodeEventProcessor: container_handlers = _container_handlers( frame_registry=frame_registry, ) - return EventHandler( + return NodeEventProcessor( graph_execution=cast(Any, graph_execution), - event_collector=cast(EventManager, event_collector), + event_stream=cast(EventStream, event_stream), frame_registry=frame_registry, container_handlers=container_handlers, ) -def _event_handler_with_container( +def _event_processor_with_container( *, graph_execution: object, - event_collector: object, + event_stream: object, frame_registry: FrameRegistry, -) -> tuple[EventHandler, dict[str, ContainerHandler]]: +) -> tuple[NodeEventProcessor, dict[str, ContainerHandler]]: container_handlers = _container_handlers( frame_registry=frame_registry, ) return ( - EventHandler( + NodeEventProcessor( graph_execution=cast(Any, graph_execution), - event_collector=cast(EventManager, event_collector), + event_stream=cast(EventStream, event_stream), frame_registry=frame_registry, container_handlers=container_handlers, ), @@ -173,7 +178,7 @@ def _get_resume_task(ready_queue: InMemoryReadyQueue) -> ResumeTask: def _start_iteration_await( container_handler: ContainerHandler, - runtime_state: GraphRuntimeState, + runtime_state: RuntimeState, *, invocation_id: str, indexes: tuple[int, ...], @@ -200,7 +205,7 @@ def _start_iteration_await( request=request, ), ) - container_handler.start_await( + container_handler.handle_request( invocation_id=invocation_id, request=request, ) @@ -209,14 +214,14 @@ def _start_iteration_await( def _worker( *, ready_queue: InMemoryReadyQueue, - event_queue: queue.Queue[DispatchTask], + dispatch_queue: queue.Queue[DispatchTask], frame_registry: FrameRegistry, ) -> Worker: task_claiming = threading.Event() task_claiming.set() return Worker( ready_queue=ready_queue, - event_queue=event_queue, + dispatch_queue=dispatch_queue, frame_registry=frame_registry, layers=[], task_claim_lock=threading.Lock(), @@ -238,7 +243,7 @@ class _FrameNode: class _FrameFactory: def with_runtime_state( self, - graph_runtime_state: GraphRuntimeState, + graph_runtime_state: RuntimeState, ) -> "_FrameFactory": _ = graph_runtime_state return self @@ -255,20 +260,20 @@ def create_node(self, node_config: dict[str, object]) -> _FrameNode: ("payload", "expected_command_type"), [ ( - {"command_type": CommandType.ABORT.value, "reason": "stop"}, + {"command_type": "abort", "reason": "stop"}, AbortCommand, ), ( - {"command_type": CommandType.PAUSE.value, "reason": "wait"}, + {"command_type": "pause", "reason": "wait"}, PauseCommand, ), ( - {"command_type": CommandType.UPDATE_VARIABLES.value, "updates": []}, + {"command_type": "update_variables", "updates": []}, UpdateVariablesCommand, ), ], ) -def test_redis_channel_deserializes_command_with_model_map( +def test_redis_channel_deserializes_discriminated_command( payload: dict[str, object], expected_command_type: type, ) -> None: @@ -279,9 +284,157 @@ def test_redis_channel_deserializes_command_with_model_map( assert isinstance(command, expected_command_type) -def test_graph_state_manager_enqueues_ready_task_for_frame() -> None: +def test_redis_channel_fetches_command_list_without_pending_marker() -> None: + """Verify one Redis transaction both reads and clears the command list. + + Fetching must not consult a second marker key: the list transaction is the + complete source of truth, and its three pipeline calls make that contract + observable without depending on a concrete Redis client implementation. + """ + redis_client = MagicMock() + pipeline = redis_client.pipeline.return_value.__enter__.return_value + pipeline.execute.return_value = [ + [AbortCommand(reason="stop").model_dump_json()], + 1, + ] + channel = RedisChannel(redis_client=redis_client, channel_key="test-channel") + + commands = channel.fetch_commands() + + assert commands == [AbortCommand(reason="stop")] + assert pipeline.method_calls == [ + call.lrange("test-channel", 0, -1), + call.delete("test-channel"), + call.execute(), + ] + + +def test_redis_channel_skips_non_object_json_and_continues_batch() -> None: + """Reject one malformed command without discarding later valid commands. + + Redis returns the entire drained list as one batch. A syntactically valid JSON + value can still be an invalid command shape, so deserialization must isolate the + bad item and continue processing the remaining entries. + """ + redis_client = MagicMock() + pipeline = redis_client.pipeline.return_value.__enter__.return_value + valid_command = AbortCommand(reason="stop") + pipeline.execute.return_value = [ + ["[]", valid_command.model_dump_json()], + 1, + ] + channel = RedisChannel(redis_client=redis_client, channel_key="test-channel") + + assert channel.fetch_commands() == [valid_command] + + +def test_redis_channel_migrates_wrapped_variable_updates() -> None: + """Accept the previous Redis payload during its bounded one-hour lifetime. + + Compatibility belongs only at the transport boundary: decoded commands + expose variables directly without restoring the removed wrapper type. + """ + variable = StringVariable( + name="answer", + selector=["node", "answer"], + value="updated", + ) + channel = RedisChannel(redis_client=MagicMock(), channel_key="test-channel") + + command = channel.deserialize_command({ + "command_type": "update_variables", + "updates": [{"value": variable.model_dump(mode="json")}], + }) + + assert isinstance(command, UpdateVariablesCommand) + assert list(command.updates) == [variable] + + +def test_command_processor_directly_handles_builtin_commands() -> None: + """Verify direct command matching preserves every built-in command behavior. + + The processor must skip an invalid variable update without dropping the valid + update that follows it, then apply pause and abort commands from the same + channel batch. This is the behavior previously split across handler classes. + """ + channel = InMemoryChannel() + execution = GraphExecution(workflow_id="workflow") + variable_pool = VariablePool() + channel.send_command( + UpdateVariablesCommand( + updates=[ + StringVariable( + name="invalid", + selector=["invalid"], + value="ignored", + ), + StringVariable( + name="answer", + selector=["node", "answer"], + value="updated", + ), + ], + ), + ) + channel.send_command(PauseCommand(reason="wait")) + channel.send_command(AbortCommand(reason="stop")) + + CommandProcessor( + command_channel=channel, + graph_execution=execution, + variable_pool=variable_pool, + ).process_commands() + + updated = variable_pool.get(["node", "answer"]) + assert updated is not None + assert updated.to_object() == "updated" + assert execution.paused + pause_reason = execution.pause_reasons[0] + assert isinstance(pause_reason, SchedulingPause) + assert pause_reason.message == "wait" + assert execution.aborted + assert str(execution.error) == "Aborted: stop" + + +def test_engine_rejects_zero_workers_before_mutating_runtime_state() -> None: + """Reject an engine that could never consume tasks before attaching its graph. + + A zero-worker engine would leave the dispatcher waiting forever. Validation + therefore belongs at the public constructor boundary and must run before the + supplied runtime state is modified. + """ + runtime_state = MagicMock() + + with pytest.raises(ValueError, match="workers must be a positive integer"): + Engine( + graph=MagicMock(), + graph_runtime_state=runtime_state, + workers=0, + ) + + runtime_state.attach_graph.assert_not_called() + + +def test_worker_pool_rejects_zero_workers() -> None: + """Reject a directly constructed pool that could never claim queued work. + + Direct callers bypass ``Engine`` validation, so the pool must enforce + the same positive-worker invariant before retaining any collaborators. + """ + with pytest.raises(ValueError, match="workers must be a positive integer"): + WorkerPool( + ready_queue=MagicMock(), + dispatch_queue=MagicMock(), + frame_registry=MagicMock(), + layers=[], + workers=0, + ) + + +def test_scheduler_enqueues_ready_task_for_frame() -> None: ready_queue = InMemoryReadyQueue() - runtime_state = GraphRuntimeState( + runtime_state = RuntimeState( + workflow_id="workflow", variable_pool=VariablePool(), start_at=0, ready_queue=ready_queue, @@ -289,26 +442,27 @@ def test_graph_state_manager_enqueues_ready_task_for_frame() -> None: graph = SimpleNamespace( nodes={"start": SimpleNamespace(state=NodeState.UNKNOWN)}, ) - manager = GraphStateManager( + scheduler = Scheduler( graph=cast(Graph, graph), - graph_runtime_state=runtime_state, + state=runtime_state, frame_id="root", ) - manager.enqueue_node("start") + scheduler.enqueue_node("start") assert ready_queue.get(timeout=0.01) == StartTask( frame_id="root", node_id="start", ) assert graph.nodes["start"].state == NodeState.TAKEN - assert not manager.is_execution_complete() + assert not scheduler.is_execution_complete() -def test_graph_state_manager_defers_ready_task_when_paused() -> None: +def test_scheduler_defers_ready_task_when_paused() -> None: ready_queue = InMemoryReadyQueue() graph_execution = GraphExecution(workflow_id="workflow", paused=True) - runtime_state = GraphRuntimeState( + runtime_state = RuntimeState( + workflow_id="workflow", variable_pool=VariablePool(), start_at=0, ready_queue=ready_queue, @@ -317,44 +471,46 @@ def test_graph_state_manager_defers_ready_task_when_paused() -> None: graph = SimpleNamespace( nodes={"start": SimpleNamespace(state=NodeState.UNKNOWN)}, ) - manager = GraphStateManager( + scheduler = Scheduler( graph=cast(Graph, graph), - graph_runtime_state=runtime_state, + state=runtime_state, frame_id="root", ) - manager.enqueue_node("start") + scheduler.enqueue_node("start") assert ready_queue.qsize() == 0 assert runtime_state.drain_deferred_ready_tasks() == [ StartTask(frame_id="root", node_id="start"), ] assert graph.nodes["start"].state == NodeState.TAKEN - assert not manager.is_execution_complete() + assert not scheduler.is_execution_complete() -def test_graph_state_manager_completion_ignores_other_frame_queue_items() -> None: +def test_scheduler_completion_ignores_other_frame_queue_items() -> None: ready_queue = InMemoryReadyQueue() ready_queue.put(StartTask(frame_id="other-frame", node_id="answer")) - runtime_state = GraphRuntimeState( + runtime_state = RuntimeState( + workflow_id="workflow", variable_pool=VariablePool(), start_at=0, ready_queue=ready_queue, ) graph = SimpleNamespace(nodes={}) - manager = GraphStateManager( + scheduler = Scheduler( graph=cast(Graph, graph), - graph_runtime_state=runtime_state, + state=runtime_state, frame_id="root", ) - assert manager.is_execution_complete() is True + assert scheduler.is_execution_complete() is True def test_pause_defers_queued_tasks_without_losing_frame_progress() -> None: ready_queue = InMemoryReadyQueue() graph_execution = GraphExecution(workflow_id="workflow") - runtime_state = GraphRuntimeState( + runtime_state = RuntimeState( + workflow_id="workflow", variable_pool=VariablePool(), start_at=0, ready_queue=ready_queue, @@ -366,13 +522,13 @@ def test_pause_defers_queued_tasks_without_losing_frame_progress() -> None: "queued": SimpleNamespace(state=NodeState.UNKNOWN), } ) - manager = GraphStateManager( + scheduler = Scheduler( graph=cast(Graph, graph), - graph_runtime_state=runtime_state, + state=runtime_state, frame_id="root", ) - manager.enqueue_node("active") - manager.enqueue_node("queued") + scheduler.enqueue_node("active") + scheduler.enqueue_node("queued") assert ready_queue.get(timeout=0.01) == StartTask( frame_id="root", node_id="active", @@ -381,13 +537,13 @@ def test_pause_defers_queued_tasks_without_losing_frame_progress() -> None: worker_pool = MagicMock() worker_pool.drain.side_effect = ready_queue.drain dispatcher = Dispatcher( - event_queue=queue.Queue(), - event_handler=MagicMock(), + dispatch_queue=queue.Queue(), + event_processor=MagicMock(), graph_execution=graph_execution, - state_manager=manager, + scheduler=scheduler, command_processor=MagicMock(), worker_pool=worker_pool, - event_emitter=MagicMock(), + event_stream=MagicMock(), ) assert dispatcher._run_until_exit() @@ -396,11 +552,11 @@ def test_pause_defers_queued_tasks_without_losing_frame_progress() -> None: assert runtime_state.drain_deferred_ready_tasks() == [ StartTask(frame_id="root", node_id="queued") ] - assert not manager.is_execution_complete() - manager.finish_execution("active") - assert not manager.is_execution_complete() - manager.finish_execution("queued") - assert manager.is_execution_complete() + assert not scheduler.is_execution_complete() + scheduler.finish_execution("active") + assert not scheduler.is_execution_complete() + scheduler.finish_execution("queued") + assert scheduler.is_execution_complete() worker_pool.drain.assert_called_once_with() worker_pool.stop.assert_not_called() @@ -409,7 +565,6 @@ def test_worker_pool_drain_does_not_stop_worker_with_current_task() -> None: class WorkerStub: def __init__(self, *, has_current_task: bool) -> None: self.has_current_task = has_current_task - self.is_idle = True self.stopped = False def stop(self) -> None: @@ -423,7 +578,6 @@ def stop(self) -> None: pool._task_claiming = threading.Event() pool._task_claiming.set() pool._ready_queue = InMemoryReadyQueue() - pool._running = True pool._workers = cast(Any, [active_worker, idle_worker]) pool.drain() @@ -493,8 +647,9 @@ def run(self) -> Generator[NodeRunSucceededEvent, None, None]: ready_queue.put(StartTask(frame_id="root", node_id="node")) node_started = threading.Event() finish_node = threading.Event() - event_queue: queue.Queue[DispatchTask] = queue.Queue() - runtime_state = GraphRuntimeState( + dispatch_queue: queue.Queue[DispatchTask] = queue.Queue() + runtime_state = RuntimeState( + workflow_id="workflow", variable_pool=VariablePool(), start_at=1, ready_queue=ready_queue, @@ -510,10 +665,10 @@ def run(self) -> Generator[NodeRunSucceededEvent, None, None]: ) pool = WorkerPool( ready_queue=ready_queue, - event_queue=event_queue, + dispatch_queue=dispatch_queue, frame_registry=frame_registry, layers=[], - config=GraphEngineConfig(max_workers=1), + workers=1, ) drained_tasks: list[ReadyTask] = [] drain_done = threading.Event() @@ -540,7 +695,7 @@ def drain_pool() -> None: pool.stop() -def test_worker_pool_scales_for_one_queued_sibling() -> None: +def test_worker_pool_runs_queued_siblings_with_fixed_workers() -> None: class ParallelNode: node_type = BuiltinNodeTypes.CODE execution_type = NodeExecutionType.EXECUTABLE @@ -568,7 +723,7 @@ def run(self) -> Generator[NodeRunSucceededEvent, None, None]: ready_queue = InMemoryReadyQueue() ready_queue.put(StartTask(frame_id="root", node_id="first")) ready_queue.put(StartTask(frame_id="root", node_id="second")) - event_queue: queue.Queue[DispatchTask] = queue.Queue() + dispatch_queue: queue.Queue[DispatchTask] = queue.Queue() first_node_started = threading.Event() barrier = threading.Barrier(2) graph = SimpleNamespace( @@ -577,7 +732,8 @@ def run(self) -> Generator[NodeRunSucceededEvent, None, None]: "second": ParallelNode("second"), }, ) - runtime_state = GraphRuntimeState( + runtime_state = RuntimeState( + workflow_id="workflow", variable_pool=VariablePool(), start_at=1, ready_queue=ready_queue, @@ -593,54 +749,46 @@ def run(self) -> Generator[NodeRunSucceededEvent, None, None]: ) pool = WorkerPool( ready_queue=ready_queue, - event_queue=event_queue, + dispatch_queue=dispatch_queue, frame_registry=frame_registry, layers=[], - config=GraphEngineConfig(max_workers=2), + workers=2, ) pool.start() try: assert first_node_started.wait(timeout=1) - pool.check_and_scale() - first_event = event_queue.get(timeout=1) - second_event = event_queue.get(timeout=1) + first_event = dispatch_queue.get(timeout=1) + second_event = dispatch_queue.get(timeout=1) finally: pool.stop() - assert isinstance(first_event, TaskEvent) - assert isinstance(second_event, TaskEvent) + assert isinstance(first_event, NodeEventTask) + assert isinstance(second_event, NodeEventTask) events = (first_event, second_event) assert {event.event.node_id for event in events} == {"first", "second"} assert all(isinstance(event.event, NodeRunSucceededEvent) for event in events) -def test_worker_with_current_task_is_not_idle() -> None: - worker = object.__new__(Worker) - worker._has_current_task = threading.Event() - worker._has_current_task.set() - worker._last_task_time = 0 - - assert not worker.is_idle - - def test_pause_requested_event_defers_current_task_for_resume() -> None: ready_queue = InMemoryReadyQueue() graph_execution = GraphExecution(workflow_id="workflow") - root_runtime_state = GraphRuntimeState( + root_runtime_state = RuntimeState( + workflow_id="workflow", variable_pool=VariablePool(), start_at=0, ready_queue=ready_queue, graph_execution=graph_execution, ) - child_runtime_state = GraphRuntimeState( + child_runtime_state = RuntimeState( + workflow_id="workflow", variable_pool=VariablePool(), start_at=0, ready_queue=ready_queue, deferred_ready_queue=root_runtime_state.deferred_ready_queue, graph_execution=graph_execution, ) - graph_init_params = GraphInitParams( + graph_init_params = InitParams( workflow_id="workflow", graph_config={}, run_context={}, @@ -708,13 +856,13 @@ def test_pause_requested_event_defers_current_task_for_resume() -> None: runtime_data=child_runtime_state.snapshot_frame(), ) ) - state_manager = GraphStateManager( + scheduler = Scheduler( graph=child_graph, - graph_runtime_state=child_runtime_state, + state=child_runtime_state, frame_id="child-frame", ) frame_registry = FrameRegistry() - event_collector = MagicMock() + event_stream = MagicMock() frame_registry.register( _execution_frame( frame_id="root", @@ -727,18 +875,18 @@ def test_pause_requested_event_defers_current_task_for_resume() -> None: frame_id="child-frame", graph=child_graph, graph_runtime_state=child_runtime_state, - state_manager=state_manager, + scheduler=scheduler, ), ) - state_manager.track_unfinished("human") - handler = _event_handler( + scheduler.track_unfinished("human") + handler = _event_processor( graph_execution=graph_execution, - event_collector=cast(EventManager, event_collector), + event_stream=cast(EventStream, event_stream), frame_registry=frame_registry, ) handler.dispatch( - TaskEvent( + NodeEventTask( frame_id="child-frame", event=NodeRunPauseRequestedEvent( id="human-run", @@ -754,7 +902,7 @@ def test_pause_requested_event_defers_current_task_for_resume() -> None: ) assert graph_execution.paused - assert not state_manager.is_execution_complete() + assert not scheduler.is_execution_complete() assert root_runtime_state.drain_deferred_ready_tasks() == [ StartTask(frame_id="child-frame", node_id="human") ] @@ -787,11 +935,11 @@ def test_graph_execution_tracks_node_executions_by_frame() -> None: assert first.execution_id != second.execution_id -def test_frame_registry_materializes_child_frame_with_rebound_runtime() -> None: +def test_frame_registry_creates_child_frame_with_rebound_runtime() -> None: @dataclass class RuntimeBoundNode: id: str - graph_runtime_state: GraphRuntimeState + graph_runtime_state: RuntimeState node_type: ClassVar[NodeType] = BuiltinNodeTypes.START execution_type: ClassVar[NodeExecutionType] = NodeExecutionType.ROOT @@ -799,12 +947,12 @@ class RuntimeBoundNode: state: ClassVar[NodeState] = NodeState.UNKNOWN class RuntimeBoundFactory: - def __init__(self, runtime_state: GraphRuntimeState) -> None: + def __init__(self, runtime_state: RuntimeState) -> None: self.runtime_state = runtime_state def with_runtime_state( self, - graph_runtime_state: GraphRuntimeState, + graph_runtime_state: RuntimeState, ) -> "RuntimeBoundFactory": return RuntimeBoundFactory(graph_runtime_state) @@ -817,7 +965,8 @@ def create_node(self, node_config: dict[str, object]) -> RuntimeBoundNode: } ready_queue = InMemoryReadyQueue() graph_execution = GraphExecution(workflow_id="workflow") - root_runtime_state = GraphRuntimeState( + root_runtime_state = RuntimeState( + workflow_id="workflow", variable_pool=VariablePool(), start_at=1, ready_queue=ready_queue, @@ -836,25 +985,22 @@ def create_node(self, node_config: dict[str, object]) -> RuntimeBoundNode: graph_runtime_state=root_runtime_state, ), ) - child_runtime_state = GraphRuntimeState( - variable_pool=VariablePool(), - start_at=2, - ready_queue=ready_queue, - graph_execution=graph_execution, - ) - - child_frame = frame_registry.materialize_child_frame( + child_frame = frame_registry.create_child( frame_id="child", + parent_frame_id="root", + container_id="", root_node_id="start", - graph_runtime_state=child_runtime_state, + variable_pool=VariablePool(), ) assert child_frame.graph is not root_graph assert child_frame.graph.nodes["start"] is not root_graph.nodes["start"] - assert child_frame.graph.nodes["start"].graph_runtime_state is child_runtime_state + assert child_frame.graph.nodes["start"].graph_runtime_state is child_frame.state + assert child_frame.state.ready_queue is root_runtime_state.ready_queue + assert child_frame.state.graph_execution is graph_execution -def test_frame_registry_materializes_child_frame_from_state() -> None: +def test_frame_registry_restores_child_frame() -> None: graph_config = { "nodes": [ {"id": "start", "data": {"type": BuiltinNodeTypes.ITERATION_START}}, @@ -863,7 +1009,8 @@ def test_frame_registry_materializes_child_frame_from_state() -> None: } ready_queue = InMemoryReadyQueue() graph_execution = GraphExecution(workflow_id="workflow") - root_runtime_state = GraphRuntimeState( + root_runtime_state = RuntimeState( + workflow_id="workflow", variable_pool=VariablePool(), start_at=1, ready_queue=ready_queue, @@ -900,8 +1047,12 @@ def test_frame_registry_materializes_child_frame_from_state() -> None: ), ) - child_frame = frame_registry.materialize_child_frame_from_state( - frame_state, + child_frame = frame_registry.restore_child( + frame_id=frame_state.frame_id, + parent_frame_id="root", + container_id="", + root_node_id=frame_state.root_node_id, + runtime_data=frame_state.runtime_data, variable_pool=cast( VariablePool, frame_state.runtime_data.variable_pool, @@ -909,11 +1060,9 @@ def test_frame_registry_materializes_child_frame_from_state() -> None: ) assert child_frame.frame_id == "child-frame" - assert _variable_value(child_frame.graph_runtime_state, ["child", "value"]) == ( - "saved" - ) - assert child_frame.graph_runtime_state.outputs == {"answer": "saved"} - assert child_frame.graph_runtime_state.node_run_steps == 2 + assert _variable_value(child_frame.state, ["child", "value"]) == ("saved") + assert child_frame.state.outputs == {"answer": "saved"} + assert child_frame.state.node_run_steps == 2 assert child_frame.graph.nodes["start"].state == NodeState.TAKEN graph_execution.pause( @@ -924,7 +1073,7 @@ def test_frame_registry_materializes_child_frame_from_state() -> None: ) ) deferred_task = StartTask(frame_id="child-frame", node_id="start") - child_frame.graph_runtime_state.enqueue_ready_task(deferred_task) + child_frame.state.enqueue_ready_task(deferred_task) assert root_runtime_state.drain_deferred_ready_tasks() == [deferred_task] @@ -937,7 +1086,8 @@ def test_frame_registry_rejects_frame_state_with_missing_graph_state_ids() -> No } ready_queue = InMemoryReadyQueue() graph_execution = GraphExecution(workflow_id="workflow") - root_runtime_state = GraphRuntimeState( + root_runtime_state = RuntimeState( + workflow_id="workflow", variable_pool=VariablePool(), start_at=1, ready_queue=ready_queue, @@ -972,16 +1122,23 @@ def test_frame_registry_rejects_frame_state_with_missing_graph_state_ids() -> No ), ) - with pytest.raises(RuntimeError, match=r"missing-node.*missing-edge"): - frame_registry.materialize_child_frame_from_state( - frame_state, + with pytest.raises( + RuntimeError, + match="Saved graph state does not match rebuilt graph", + ): + frame_registry.restore_child( + frame_id=frame_state.frame_id, + parent_frame_id="root", + container_id="", + root_node_id=frame_state.root_node_id, + runtime_data=frame_state.runtime_data, variable_pool=cast( VariablePool, frame_state.runtime_data.variable_pool, ).model_copy(deep=True), ) with pytest.raises(KeyError): - frame_registry.get("child-frame") + frame_registry["child-frame"] def test_frame_registry_copies_frame_runtime_data_from_state() -> None: @@ -993,7 +1150,8 @@ def test_frame_registry_copies_frame_runtime_data_from_state() -> None: } ready_queue = InMemoryReadyQueue() graph_execution = GraphExecution(workflow_id="workflow") - root_runtime_state = GraphRuntimeState( + root_runtime_state = RuntimeState( + workflow_id="workflow", variable_pool=VariablePool(), start_at=1, ready_queue=ready_queue, @@ -1025,20 +1183,24 @@ def test_frame_registry_copies_frame_runtime_data_from_state() -> None: outputs={"nested": {"value": "saved"}}, llm_usage=LLMUsage.empty_usage(), node_run_steps=0, - graph_node_states={}, + graph_node_states={"start": NodeState.UNKNOWN}, graph_edge_states={}, ), ) - child_frame = frame_registry.materialize_child_frame_from_state( - frame_state, + child_frame = frame_registry.restore_child( + frame_id=frame_state.frame_id, + parent_frame_id="root", + container_id="", + root_node_id=frame_state.root_node_id, + runtime_data=frame_state.runtime_data, variable_pool=cast( VariablePool, frame_state.runtime_data.variable_pool, ).model_copy(deep=True), ) - child_frame.graph_runtime_state.variable_pool.add(["child", "value"], "changed") - child_frame.graph_runtime_state.set_output("nested", {"value": "changed"}) + child_frame.state.variable_pool.add(["child", "value"], "changed") + child_frame.state.set_output("nested", {"value": "changed"}) saved_variable = cast( VariablePool, @@ -1070,9 +1232,10 @@ def run(self) -> Generator[NodeRunStartedEvent, None, None]: ready_queue = InMemoryReadyQueue() ready_queue.put(StartTask(frame_id="root", node_id="start")) - event_queue = queue.Queue() + dispatch_queue = queue.Queue() graph = SimpleNamespace(nodes={"start": RunnableNode()}) - runtime_state = GraphRuntimeState( + runtime_state = RuntimeState( + workflow_id="workflow", variable_pool=VariablePool(), start_at=1, ready_queue=ready_queue, @@ -1088,18 +1251,18 @@ def run(self) -> Generator[NodeRunStartedEvent, None, None]: ) worker = _worker( ready_queue=ready_queue, - event_queue=event_queue, + dispatch_queue=dispatch_queue, frame_registry=frame_registry, ) worker.start() try: - event = event_queue.get(timeout=1) + event = dispatch_queue.get(timeout=1) finally: worker.stop() worker.join(timeout=1) - assert isinstance(event, TaskEvent) + assert isinstance(event, NodeEventTask) assert event.frame_id == "root" assert isinstance(event.event, NodeRunStartedEvent) assert event.event.node_id == "start" @@ -1129,11 +1292,12 @@ def run(self) -> Generator[NodeRunStartedEvent, None, None]: ready_queue = InMemoryReadyQueue() ready_queue.put(StartTask(frame_id="child", node_id="answer")) - event_queue = queue.Queue() + dispatch_queue = queue.Queue() root_graph = SimpleNamespace(nodes={"answer": RunnableNode("answer", "Root")}) child_graph = SimpleNamespace(nodes={"answer": RunnableNode("answer", "Child")}) graph_execution = GraphExecution(workflow_id="workflow") - runtime_state = GraphRuntimeState( + runtime_state = RuntimeState( + workflow_id="workflow", variable_pool=VariablePool(), start_at=1, ready_queue=ready_queue, @@ -1156,18 +1320,18 @@ def run(self) -> Generator[NodeRunStartedEvent, None, None]: ) worker = _worker( ready_queue=ready_queue, - event_queue=event_queue, + dispatch_queue=dispatch_queue, frame_registry=frame_registry, ) worker.start() try: - event = event_queue.get(timeout=1) + event = dispatch_queue.get(timeout=1) finally: worker.stop() worker.join(timeout=1) - assert isinstance(event, TaskEvent) + assert isinstance(event, NodeEventTask) assert event.frame_id == "child" assert isinstance(event.event, NodeRunStartedEvent) assert event.event.node_title == "Child" @@ -1196,7 +1360,7 @@ def run(self) -> Generator[NodeRunStartedEvent, None, None]: ready_queue = InMemoryReadyQueue() ready_queue.put(StartTask(frame_id="child", node_id="answer")) - event_queue = queue.Queue() + dispatch_queue = queue.Queue() graph_execution = GraphExecution(workflow_id="workflow") root_execution = graph_execution.get_or_create_node_execution( frame_id="root", @@ -1206,7 +1370,8 @@ def run(self) -> Generator[NodeRunStartedEvent, None, None]: frame_id="child", node_id="answer", ) - runtime_state = GraphRuntimeState( + runtime_state = RuntimeState( + workflow_id="workflow", variable_pool=VariablePool(), start_at=1, ready_queue=ready_queue, @@ -1223,18 +1388,18 @@ def run(self) -> Generator[NodeRunStartedEvent, None, None]: ) worker = _worker( ready_queue=ready_queue, - event_queue=event_queue, + dispatch_queue=dispatch_queue, frame_registry=frame_registry, ) worker.start() try: - event = event_queue.get(timeout=1) + event = dispatch_queue.get(timeout=1) finally: worker.stop() worker.join(timeout=1) - assert isinstance(event, TaskEvent) + assert isinstance(event, NodeEventTask) assert isinstance(event.event, NodeRunStartedEvent) assert event.event.id != root_execution.execution_id assert event.event.id == child_execution.execution_id @@ -1248,11 +1413,11 @@ def test_dispatcher_preserves_task_event_for_dispatch() -> None: node_title="Start", start_at=datetime.now(UTC).replace(tzinfo=None), ) - task_event = TaskEvent(frame_id="root", event=event) - event_queue = queue.Queue() - event_queue.put(task_event) + task_event = NodeEventTask(frame_id="root", event=event) + dispatch_queue = queue.Queue() + dispatch_queue.put(task_event) - class RecordingEventHandler: + class RecordingNodeEventProcessor: dispatched_events: list[object] def __init__(self) -> None: @@ -1261,33 +1426,33 @@ def __init__(self) -> None: def dispatch(self, event: object) -> None: self.dispatched_events.append(event) - event_handler = RecordingEventHandler() + event_processor = RecordingNodeEventProcessor() graph_execution = MagicMock( aborted=False, paused=False, error=None, completed=False, ) - state_manager = MagicMock() - state_manager.is_execution_complete.side_effect = lambda: bool( - event_handler.dispatched_events + scheduler = MagicMock() + scheduler.is_execution_complete.side_effect = lambda: bool( + event_processor.dispatched_events ) dispatcher = Dispatcher( - event_queue=event_queue, - event_handler=cast(EventHandler, event_handler), + dispatch_queue=dispatch_queue, + event_processor=cast(NodeEventProcessor, event_processor), graph_execution=graph_execution, - state_manager=state_manager, + scheduler=scheduler, command_processor=MagicMock(), worker_pool=MagicMock(), - event_emitter=MagicMock(), + event_stream=MagicMock(), ) dispatcher._dispatcher_loop() - assert event_handler.dispatched_events == [task_event] + assert event_processor.dispatched_events == [task_event] -def test_event_handler_dispatches_task_event_payload() -> None: +def test_event_processor_dispatches_task_event_payload() -> None: event = NodeRunStartedEvent( id="run-1", node_id="start", @@ -1298,7 +1463,7 @@ def test_event_handler_dispatches_task_event_payload() -> None: node_execution = MagicMock(retry_count=0) graph_execution = MagicMock() graph_execution.get_or_create_node_execution.return_value = node_execution - event_collector = MagicMock() + event_stream = MagicMock() runtime_state = MagicMock() frame_registry = FrameRegistry() frame_registry.register( @@ -1308,42 +1473,51 @@ def test_event_handler_dispatches_task_event_payload() -> None: graph_runtime_state=runtime_state, ), ) - handler = _event_handler( + handler = _event_processor( graph_execution=graph_execution, - event_collector=cast(EventManager, event_collector), + event_stream=cast(EventStream, event_stream), frame_registry=frame_registry, ) - handler.dispatch(TaskEvent(frame_id="root", event=event)) + handler.dispatch(NodeEventTask(frame_id="root", event=event)) runtime_state.increment_node_run_steps.assert_called_once_with() - event_collector.collect.assert_called_once_with(event) + event_stream.collect.assert_called_once_with(event) -def test_event_handler_processes_tagged_root_frame_success_before_collecting() -> None: +def test_event_processor_stamps_frame_owner_on_node_and_edge_events() -> None: graph = MagicMock() graph.nodes = {"child": MagicMock(execution_type=NodeExecutionType.EXECUTABLE)} runtime_state = MagicMock() runtime_state.variable_pool = MagicMock() graph_execution = MagicMock() graph_execution.get_or_create_node_execution.return_value = MagicMock() - event_collector = MagicMock() - edge_processor = MagicMock() - edge_processor.process_node_success.return_value = (["next"], []) - state_manager = MagicMock() + event_stream = MagicMock() + scheduler = MagicMock() + taken = GraphEdgeTakenEvent( + edge_id="taken", + source_node_id="child", + target_node_id="next", + ) + skipped = GraphEdgeSkippedEvent( + edge_id="skipped", + source_node_id="child", + target_node_id="other", + ) + scheduler.process_node_success.return_value = (["next"], [taken, skipped]) frame_registry = FrameRegistry() frame_registry.register( _execution_frame( frame_id="root", graph=cast(Graph, graph), graph_runtime_state=runtime_state, - state_manager=state_manager, - edge_processor=edge_processor, + scheduler=scheduler, + container_id="owner", ), ) - handler = _event_handler( + handler = _event_processor( graph_execution=graph_execution, - event_collector=cast(EventManager, event_collector), + event_stream=cast(EventStream, event_stream), frame_registry=frame_registry, ) event = NodeRunSucceededEvent( @@ -1353,14 +1527,21 @@ def test_event_handler_processes_tagged_root_frame_success_before_collecting() - start_at=datetime.now(UTC).replace(tzinfo=None), finished_at=datetime.now(UTC).replace(tzinfo=None), node_run_result=NodeRunResult(outputs={"answer": "ok"}), - in_iteration_id="iteration", + container_id="stale", ) - handler.dispatch(TaskEvent(frame_id="root", event=event)) + handler.dispatch(NodeEventTask(frame_id="root", event=event)) - edge_processor.process_node_success.assert_called_once_with("child") - state_manager.enqueue_node.assert_called_once_with("next") - event_collector.collect.assert_called_once_with(event) + scheduler.process_node_success.assert_called_once_with("child") + scheduler.enqueue_node.assert_called_once_with("next") + assert event.container_id == "owner" + assert taken.container_id == "owner" + assert skipped.container_id == "owner" + assert event_stream.collect.call_args_list == [ + call(taken), + call(skipped), + call(event), + ] def test_parallel_iteration_preserves_aggregate_and_response_order() -> None: # ruff: ignore[too-many-locals] @@ -1368,7 +1549,10 @@ def test_parallel_iteration_preserves_aggregate_and_response_order() -> None: # "nodes": [ { "id": "iteration-start", - "data": {"type": BuiltinNodeTypes.ITERATION_START}, + "data": { + "type": BuiltinNodeTypes.ITERATION_START, + "container_id": "iteration", + }, }, ], "edges": [], @@ -1377,7 +1561,8 @@ def test_parallel_iteration_preserves_aggregate_and_response_order() -> None: # variable_pool = VariablePool() variable_pool.add(["source", "items"], ["a", "b", "c"]) graph_execution = GraphExecution(workflow_id="workflow") - runtime_state = GraphRuntimeState( + runtime_state = RuntimeState( + workflow_id="workflow", variable_pool=variable_pool, start_at=1, ready_queue=ready_queue, @@ -1409,10 +1594,10 @@ def test_parallel_iteration_preserves_aggregate_and_response_order() -> None: # graph_runtime_state=runtime_state, ), ) - event_collector = MagicMock() - handler, container_handlers = _event_handler_with_container( + event_stream = MagicMock() + handler, container_handlers = _event_processor_with_container( graph_execution=graph_execution, - event_collector=cast(EventManager, event_collector), + event_stream=cast(EventStream, event_stream), frame_registry=frame_registry, ) @@ -1437,11 +1622,11 @@ def test_parallel_iteration_preserves_aggregate_and_response_order() -> None: # ) assert ready_queue.qsize() == 0 - second_frame = frame_registry.get("iteration-invocation:iteration:1") - second_frame.graph_runtime_state.variable_pool.add(["answer", "text"], "second") - second_frame.graph_runtime_state.set_output("answer", "second") + second_frame = frame_registry["iteration-invocation:iteration:1"] + second_frame.state.variable_pool.add(["answer", "text"], "second") + second_frame.state.set_output("answer", "second") handler.dispatch( - TaskEvent( + NodeEventTask( frame_id="iteration-invocation:iteration:1", event=NodeRunSucceededEvent( id="iteration-start-run-1", @@ -1457,7 +1642,7 @@ def test_parallel_iteration_preserves_aggregate_and_response_order() -> None: # resume_task = _get_resume_task(ready_queue) assert isinstance(resume_task.result, IterationFrameRequest) assert resume_task.result.indexes == (2,) - container_handlers["iteration"].start_await( + container_handlers["iteration"].handle_request( invocation_id=resume_task.invocation_id, request=resume_task.result, ) @@ -1466,11 +1651,11 @@ def test_parallel_iteration_preserves_aggregate_and_response_order() -> None: # node_id="iteration-start", ) - third_frame = frame_registry.get("iteration-invocation:iteration:2") - third_frame.graph_runtime_state.variable_pool.add(["answer", "text"], "third") - third_frame.graph_runtime_state.set_output("answer", "third") + third_frame = frame_registry["iteration-invocation:iteration:2"] + third_frame.state.variable_pool.add(["answer", "text"], "third") + third_frame.state.set_output("answer", "third") handler.dispatch( - TaskEvent( + NodeEventTask( frame_id="iteration-invocation:iteration:2", event=NodeRunSucceededEvent( id="iteration-start-run-2", @@ -1483,11 +1668,11 @@ def test_parallel_iteration_preserves_aggregate_and_response_order() -> None: # ), ) - first_frame = frame_registry.get("iteration-invocation:iteration:0") - first_frame.graph_runtime_state.variable_pool.add(["answer", "text"], "first") - first_frame.graph_runtime_state.set_output("answer", "first") + first_frame = frame_registry["iteration-invocation:iteration:0"] + first_frame.state.variable_pool.add(["answer", "text"], "first") + first_frame.state.set_output("answer", "first") handler.dispatch( - TaskEvent( + NodeEventTask( frame_id="iteration-invocation:iteration:0", event=NodeRunSucceededEvent( id="iteration-start-run-0", @@ -1512,7 +1697,8 @@ def test_parallel_iteration_preserves_aggregate_and_response_order() -> None: # def test_terminated_iteration_waits_for_all_scheduled_frames() -> None: ready_queue = InMemoryReadyQueue() - runtime_state = GraphRuntimeState( + runtime_state = RuntimeState( + workflow_id="workflow", variable_pool=VariablePool(), start_at=1, ready_queue=ready_queue, @@ -1534,11 +1720,11 @@ def test_terminated_iteration_waits_for_all_scheduled_frames() -> None: ) runtime_state.put_container_run(run_state) frame_registry = MagicMock() - frame_registry.get.return_value.graph_runtime_state = runtime_state + frame_registry.__getitem__.return_value.state = runtime_state handler = IterationContainerHandler(frame_registry=frame_registry) parent_frame = cast( ExecutionFrame, - SimpleNamespace(graph_runtime_state=runtime_state), + SimpleNamespace(state=runtime_state), ) assert handler._finish_failed_iteration_if_ready( @@ -1565,7 +1751,10 @@ def test_iteration_frame_completion_requests_next_index() -> None: "nodes": [ { "id": "iteration-start", - "data": {"type": BuiltinNodeTypes.ITERATION_START}, + "data": { + "type": BuiltinNodeTypes.ITERATION_START, + "container_id": "iteration", + }, }, ], "edges": [], @@ -1574,7 +1763,8 @@ def test_iteration_frame_completion_requests_next_index() -> None: variable_pool = VariablePool() variable_pool.add(["source", "items"], ["a", "b", "c"]) graph_execution = GraphExecution(workflow_id="workflow") - runtime_state = GraphRuntimeState( + runtime_state = RuntimeState( + workflow_id="workflow", variable_pool=variable_pool, start_at=1, ready_queue=ready_queue, @@ -1628,11 +1818,11 @@ def test_iteration_frame_completion_requests_next_index() -> None: node_id="iteration-start", ) - sibling_frame = frame_registry.get("iteration-invocation:iteration:1") - sibling_frame.graph_runtime_state.variable_pool.add(["answer", "text"], "second") - sibling_frame.state_manager.finish_execution("iteration-start") + sibling_frame = frame_registry["iteration-invocation:iteration:1"] + sibling_frame.state.variable_pool.add(["answer", "text"], "second") + sibling_frame.scheduler.finish_execution("iteration-start") - container_handler.complete_frame(sibling_frame) + container_handler.complete_frame_if_ready(sibling_frame) run_state = runtime_state.get_container_run("iteration-invocation") assert isinstance(run_state, IterationRunState) @@ -1644,25 +1834,41 @@ def test_iteration_frame_completion_requests_next_index() -> None: @pytest.mark.parametrize( - ("limit_type", "expected_reason"), + ("max_steps", "max_time", "step_count", "elapsed_time", "expected_reason"), [ - (LimitType.STEP_LIMIT, "Maximum execution steps exceeded: 4 > 3"), - (LimitType.TIME_LIMIT, "Maximum execution time exceeded:"), + (3, 1000, 4, 0, "Maximum execution steps exceeded: 4 > 3"), + (10, 10, 0, 20, "Maximum execution time exceeded:"), ], ) -def test_execution_limits_layer_builds_abort_reason_with_match_case( - limit_type: LimitType, +def test_execution_limits_layer_sends_abort_when_limit_is_exceeded( + max_steps: int, + max_time: int, + step_count: int, + elapsed_time: int, expected_reason: str, ) -> None: - layer = ExecutionLimitsLayer(max_steps=3, max_time=10) - layer.command_channel = MagicMock() + """Verify the event hook emits an abort for step and elapsed-time limits. + + The parameter sets isolate one exceeded limit at a time so the assertion + covers the production event path without exposing a test-only layer method. + """ + layer = ExecutionLimitsLayer(max_steps=max_steps, max_time=max_time) + command_channel = MagicMock() + layer.command_channel = command_channel layer.on_graph_start() - layer.step_count = 4 - layer.start_time = time() - 20 + layer.step_count = step_count + layer.start_time = time() - elapsed_time - layer.send_abort_command(limit_type) + layer.on_event( + NodeRunSucceededEvent( + id="node-run-1", + node_id="node-1", + node_type=BuiltinNodeTypes.CODE, + start_at=datetime.now(UTC).replace(tzinfo=None), + ) + ) - abort_command = layer.command_channel.send_command.call_args.args[0] + abort_command = command_channel.send_command.call_args.args[0] assert isinstance(abort_command, AbortCommand) assert abort_command.reason is not None assert abort_command.reason.startswith(expected_reason) diff --git a/tests/graph_engine/test_event_filters.py b/tests/engine/test_event_filters.py similarity index 58% rename from tests/graph_engine/test_event_filters.py rename to tests/engine/test_event_filters.py index 7e9802a7..b8ea05a2 100644 --- a/tests/graph_engine/test_event_filters.py +++ b/tests/engine/test_event_filters.py @@ -1,79 +1,71 @@ from collections.abc import Iterable from typing import Any, cast -from graphon.filters import ( - GraphEventFilter, - GraphEventFilterContext, - filter_graph_events, +from graphon.engine.filter import ( + EngineEventFilter, + EngineEventFilterContext, + filter_engine_events, ) -from graphon.graph_events.base import GraphEngineEvent -from graphon.graph_events.graph import GraphRunStartedEvent -from graphon.graph_events.traversal import GraphEdgeTakenEvent +from graphon.engine_events.base import EngineEvent +from graphon.engine_events.graph import GraphRunStartedEvent +from graphon.engine_events.traversal import GraphEdgeTakenEvent -def _context() -> GraphEventFilterContext: - return GraphEventFilterContext( +def _context() -> EngineEventFilterContext: + return EngineEventFilterContext( graph=cast(Any, object()), runtime_state=cast(Any, object()), ) class _PassThroughFilter: - filter_id = "pass-through" - def __init__(self) -> None: self.initialized = False - def initialize(self, context: GraphEventFilterContext) -> None: + def initialize(self, context: EngineEventFilterContext) -> None: self.initialized = context is not None - def on_event(self, event: GraphEngineEvent) -> Iterable[GraphEngineEvent]: + def on_event(self, event: EngineEvent) -> Iterable[EngineEvent]: yield event - def flush(self) -> Iterable[GraphEngineEvent]: + def flush(self) -> Iterable[EngineEvent]: return () class _DropTraversalFilter: - filter_id = "drop-traversal" - - def initialize(self, context: GraphEventFilterContext) -> None: + def initialize(self, context: EngineEventFilterContext) -> None: self.context = context - def on_event(self, event: GraphEngineEvent) -> Iterable[GraphEngineEvent]: + def on_event(self, event: EngineEvent) -> Iterable[EngineEvent]: if isinstance(event, GraphEdgeTakenEvent): return () return (event,) - def flush(self) -> Iterable[GraphEngineEvent]: + def flush(self) -> Iterable[EngineEvent]: return () class _SplitStartFilter: - filter_id = "split-start" - - def initialize(self, context: GraphEventFilterContext) -> None: + def initialize(self, context: EngineEventFilterContext) -> None: self.context = context - def on_event(self, event: GraphEngineEvent) -> Iterable[GraphEngineEvent]: + def on_event(self, event: EngineEvent) -> Iterable[EngineEvent]: if isinstance(event, GraphRunStartedEvent): return (event, event.model_copy()) return (event,) - def flush(self) -> Iterable[GraphEngineEvent]: + def flush(self) -> Iterable[EngineEvent]: return () class _FlushFilter: - filter_id = "flush" - - def initialize(self, context: GraphEventFilterContext) -> None: + def initialize(self, context: EngineEventFilterContext) -> None: self.context = context - def on_event(self, event: GraphEngineEvent) -> Iterable[GraphEngineEvent]: + def on_event(self, event: EngineEvent) -> Iterable[EngineEvent]: return (event,) - def flush(self) -> Iterable[GraphEngineEvent]: + def flush(self) -> Iterable[EngineEvent]: return ( GraphEdgeTakenEvent( edge_id="flush-edge", @@ -84,14 +76,15 @@ def flush(self) -> Iterable[GraphEngineEvent]: def test_filter_protocol_accepts_pass_through_filter() -> None: - event_filter: GraphEventFilter = _PassThroughFilter() - assert event_filter.filter_id == "pass-through" + event_filter: EngineEventFilter = _PassThroughFilter() + event_filter.initialize(_context()) + assert event_filter.initialized is True def test_filter_chain_passes_events_when_no_filters() -> None: event = GraphRunStartedEvent() output = list( - filter_graph_events( + filter_engine_events( [event], context=_context(), filters=[], @@ -111,7 +104,7 @@ def test_filter_chain_initializes_and_chains_drop_and_split() -> None: start = GraphRunStartedEvent() output = list( - filter_graph_events( + filter_engine_events( [start, edge], context=_context(), filters=[pass_through, _SplitStartFilter(), _DropTraversalFilter()], @@ -124,7 +117,7 @@ def test_filter_chain_initializes_and_chains_drop_and_split() -> None: def test_filter_chain_sends_flush_output_to_downstream_filters() -> None: output = list( - filter_graph_events( + filter_engine_events( [], context=_context(), filters=[_FlushFilter(), _DropTraversalFilter()], diff --git a/tests/graph_engine/test_raw_engine_events.py b/tests/engine/test_raw_engine_events.py similarity index 51% rename from tests/graph_engine/test_raw_engine_events.py rename to tests/engine/test_raw_engine_events.py index c24a23db..54846384 100644 --- a/tests/graph_engine/test_raw_engine_events.py +++ b/tests/engine/test_raw_engine_events.py @@ -5,17 +5,17 @@ import pytest -from graphon import graph_events, node_events -from graphon.enums import BuiltinNodeTypes, NodeExecutionType -from graphon.graph_engine.entities.tasks import TaskEvent -from graphon.graph_engine.event_management.event_handlers import EventHandler -from graphon.graph_engine.frames import ExecutionFrame, FrameRegistry -from graphon.graph_events.node import ( +from graphon import engine_events, node_events +from graphon.engine.event.processor import NodeEventProcessor +from graphon.engine.frame import ExecutionFrame, FrameRegistry +from graphon.engine.worker import NodeEventTask +from graphon.engine_events.node import ( NodeRunReasoningChunkEvent, NodeRunStreamChunkEvent, NodeRunSucceededEvent, ) -from graphon.graph_events.traversal import GraphEdgeTakenEvent +from graphon.engine_events.traversal import GraphEdgeTakenEvent +from graphon.enums import BuiltinNodeTypes, NodeExecutionType from graphon.node_events.base import NodeRunResult @@ -26,52 +26,49 @@ def _now() -> datetime: def _root_frame( *, graph: object, - graph_runtime_state: object, - state_manager: object, - edge_processor: object, - error_handler: object, + state: object, + scheduler: object, + failure_handler: object, ) -> FrameRegistry: - if isinstance(graph_runtime_state, MagicMock): - graph_runtime_state.has_container_frame.return_value = False + if isinstance(state, MagicMock): + state.has_container_frame.return_value = False frame_registry = FrameRegistry() frame_registry.register( ExecutionFrame( frame_id="root", graph=cast(Any, graph), - graph_runtime_state=cast(Any, graph_runtime_state), - state_manager=cast(Any, state_manager), - edge_processor=cast(Any, edge_processor), - error_handler=cast(Any, error_handler), + state=cast(Any, state), + scheduler=cast(Any, scheduler), + failure_handler=cast(Any, failure_handler), ), ) return frame_registry -def _event_handler( +def _event_processor( *, graph_execution: object, - event_collector: object, + event_stream: object, frame_registry: FrameRegistry, -) -> EventHandler: - return EventHandler( +) -> NodeEventProcessor: + return NodeEventProcessor( graph_execution=cast(Any, graph_execution), - event_collector=cast(Any, event_collector), + event_stream=cast(Any, event_stream), frame_registry=frame_registry, container_handlers={}, ) -def test_event_handler_collects_raw_stream_chunk_without_coordinator() -> None: - event_collector = MagicMock() - handler = _event_handler( +def test_event_processor_collects_raw_stream_chunk_without_coordinator() -> None: + event_stream = MagicMock() + handler = _event_processor( graph_execution=cast(Any, MagicMock()), - event_collector=cast(Any, event_collector), + event_stream=cast(Any, event_stream), frame_registry=_root_frame( graph=MagicMock(), - graph_runtime_state=MagicMock(), - state_manager=MagicMock(), - edge_processor=MagicMock(), - error_handler=MagicMock(), + state=MagicMock(), + scheduler=MagicMock(), + failure_handler=MagicMock(), ), ) chunk = NodeRunStreamChunkEvent( @@ -83,26 +80,25 @@ def test_event_handler_collects_raw_stream_chunk_without_coordinator() -> None: is_final=False, ) - handler.dispatch(TaskEvent(frame_id="root", event=chunk)) + handler.dispatch(NodeEventTask(frame_id="root", event=chunk)) - event_collector.collect.assert_called_once_with(chunk) + event_stream.collect.assert_called_once_with(chunk) -def test_event_handler_collects_reasoning_chunk_without_warning( +def test_event_processor_collects_reasoning_chunk_without_warning( caplog: pytest.LogCaptureFixture, ) -> None: # Reasoning chunks must hit the registered collect-only group, not the # default fallback that warns once per chunk. - event_collector = MagicMock() - handler = _event_handler( + event_stream = MagicMock() + handler = _event_processor( graph_execution=cast(Any, MagicMock()), - event_collector=cast(Any, event_collector), + event_stream=cast(Any, event_stream), frame_registry=_root_frame( graph=MagicMock(), - graph_runtime_state=MagicMock(), - state_manager=MagicMock(), - edge_processor=MagicMock(), - error_handler=MagicMock(), + state=MagicMock(), + scheduler=MagicMock(), + failure_handler=MagicMock(), ), ) chunk = NodeRunReasoningChunkEvent( @@ -115,44 +111,42 @@ def test_event_handler_collects_reasoning_chunk_without_warning( ) with caplog.at_level(logging.WARNING): - handler.dispatch(TaskEvent(frame_id="root", event=chunk)) + handler.dispatch(NodeEventTask(frame_id="root", event=chunk)) - event_collector.collect.assert_called_once_with(chunk) + event_stream.collect.assert_called_once_with(chunk) assert "Unhandled event type" not in caplog.text def test_reasoning_events_are_exported_from_package_roots() -> None: - assert graph_events.NodeRunReasoningChunkEvent is NodeRunReasoningChunkEvent - assert "NodeRunReasoningChunkEvent" in graph_events.__all__ + assert engine_events.NodeRunReasoningChunkEvent is NodeRunReasoningChunkEvent + assert "NodeRunReasoningChunkEvent" in engine_events.__all__ assert "StreamReasoningEvent" in node_events.__all__ -def test_event_handler_collects_traversal_events_before_node_success() -> None: +def test_event_processor_collects_traversal_events_before_node_success() -> None: graph = MagicMock() graph.nodes = {"node-1": MagicMock(execution_type=NodeExecutionType.EXECUTABLE)} runtime_state = MagicMock() runtime_state.variable_pool = MagicMock() graph_execution = MagicMock() graph_execution.get_or_create_node_execution.return_value = MagicMock() - event_collector = MagicMock() + event_stream = MagicMock() edge_event = GraphEdgeTakenEvent( edge_id="edge-1", source_node_id="node-1", target_node_id="node-2", source_handle="success", ) - edge_processor = MagicMock() - edge_processor.process_node_success.return_value = ([], [edge_event]) - state_manager = MagicMock() - handler = _event_handler( + scheduler = MagicMock() + scheduler.process_node_success.return_value = ([], [edge_event]) + handler = _event_processor( graph_execution=cast(Any, graph_execution), - event_collector=cast(Any, event_collector), + event_stream=cast(Any, event_stream), frame_registry=_root_frame( graph=graph, - graph_runtime_state=runtime_state, - state_manager=state_manager, - edge_processor=edge_processor, - error_handler=MagicMock(), + state=runtime_state, + scheduler=scheduler, + failure_handler=MagicMock(), ), ) success = NodeRunSucceededEvent( @@ -164,7 +158,7 @@ def test_event_handler_collects_traversal_events_before_node_success() -> None: node_run_result=NodeRunResult(outputs={"answer": "hello"}), ) - handler.dispatch(TaskEvent(frame_id="root", event=success)) + handler.dispatch(NodeEventTask(frame_id="root", event=success)) - collected_events = [call.args[0] for call in event_collector.collect.call_args_list] + collected_events = [call.args[0] for call in event_stream.collect.call_args_list] assert collected_events == [edge_event, success] diff --git a/tests/graph_engine/test_response_stream_filter.py b/tests/engine/test_response_stream_filter.py similarity index 97% rename from tests/graph_engine/test_response_stream_filter.py rename to tests/engine/test_response_stream_filter.py index 3d253d66..4860e42b 100644 --- a/tests/graph_engine/test_response_stream_filter.py +++ b/tests/engine/test_response_stream_filter.py @@ -4,20 +4,20 @@ import pytest -from graphon.enums import BuiltinNodeTypes, NodeExecutionType, NodeState, NodeType -from graphon.filters import ( - GraphEventFilterContext, +from graphon.engine.filter import ( + EngineEventFilterContext, ResponseStreamFilter, - filter_graph_events, + filter_engine_events, ) -from graphon.graph_events.graph import GraphRunStartedEvent -from graphon.graph_events.node import ( +from graphon.engine_events.graph import GraphRunStartedEvent +from graphon.engine_events.node import ( NodeRunReasoningChunkEvent, NodeRunRetryEvent, NodeRunStartedEvent, NodeRunStreamChunkEvent, ) -from graphon.graph_events.traversal import GraphEdgeTakenEvent +from graphon.engine_events.traversal import GraphEdgeTakenEvent +from graphon.enums import BuiltinNodeTypes, NodeExecutionType, NodeState, NodeType from graphon.nodes.base.template import ( Template, TemplateSegmentUnion, @@ -27,8 +27,8 @@ from graphon.runtime.graph_runtime_state import ( EdgeProtocol, GraphProtocol, - GraphRuntimeState, NodeProtocol, + RuntimeState, ) from graphon.runtime.read_only_wrappers import ReadOnlyGraphRuntimeStateWrapper from graphon.runtime.variable_pool import VariablePool @@ -121,10 +121,14 @@ def get_incoming_edges(self, node_id: str) -> Sequence[_TestEdge]: def _context( graph: _TestGraph, variable_pool: VariablePool | None = None, -) -> GraphEventFilterContext: - state = GraphRuntimeState(variable_pool=variable_pool or VariablePool(), start_at=0) +) -> EngineEventFilterContext: + state = RuntimeState( + workflow_id="workflow", + variable_pool=variable_pool or VariablePool(), + start_at=0, + ) state.attach_graph(cast(Any, graph)) - return GraphEventFilterContext( + return EngineEventFilterContext( graph=cast(Any, graph), runtime_state=ReadOnlyGraphRuntimeStateWrapper(state), ) @@ -621,7 +625,7 @@ def test_response_stream_filter_can_load_before_filter_chain_initializes() -> No restored_filter.loads(snapshot) assert restored_filter.dumps() == snapshot output = list( - filter_graph_events( + filter_engine_events( [_edge_taken()], context=context, filters=[restored_filter], @@ -641,7 +645,7 @@ def test_response_stream_filter_restores_referenced_selectors() -> None: restored_filter = ResponseStreamFilter() restored_filter.loads(first_filter.dumps()) output = list( - filter_graph_events( + filter_engine_events( [ _edge_taken(), _stream_chunk("late"), diff --git a/tests/graph_engine/graph_traversal/test_traversal_events.py b/tests/engine/test_scheduler.py similarity index 54% rename from tests/graph_engine/graph_traversal/test_traversal_events.py rename to tests/engine/test_scheduler.py index 03e8d22f..b4c8faa1 100644 --- a/tests/graph_engine/graph_traversal/test_traversal_events.py +++ b/tests/engine/test_scheduler.py @@ -1,13 +1,13 @@ from collections.abc import Sequence from typing import cast +from graphon.engine.scheduler import Scheduler +from graphon.engine_events.traversal import GraphEdgeSkippedEvent, GraphEdgeTakenEvent from graphon.enums import BuiltinNodeTypes, NodeExecutionType, NodeState from graphon.graph.edge import Edge from graphon.graph.graph import Graph -from graphon.graph_engine.graph_state_manager import GraphStateManager -from graphon.graph_engine.graph_traversal.edge_processor import EdgeProcessor -from graphon.graph_engine.graph_traversal.skip_propagator import SkipPropagator -from graphon.graph_events.traversal import GraphEdgeSkippedEvent, GraphEdgeTakenEvent +from graphon.runtime.graph_runtime_state import RuntimeState +from graphon.runtime.variable_pool import VariablePool type _TraversalEvent = GraphEdgeTakenEvent | GraphEdgeSkippedEvent @@ -57,58 +57,14 @@ def get_incoming_edges(self, node_id: str) -> list[Edge]: return [edge for edge in self.edges.values() if edge.head == node_id] -class _StateManager: - def __init__(self, graph: _Graph) -> None: - self.graph = graph - self.started_nodes: list[str] = [] - - def categorize_branch_edges( - self, - node_id: str, - selected_handle: str, - ) -> tuple[list[Edge], list[Edge]]: - edges = self.graph.get_outgoing_edges(node_id) - return ( - [edge for edge in edges if edge.source_handle == selected_handle], - [edge for edge in edges if edge.source_handle != selected_handle], - ) - - def mark_edge_taken(self, edge_id: str) -> None: - self.graph.edges[edge_id].state = NodeState.TAKEN - - def mark_edge_skipped(self, edge_id: str) -> None: - self.graph.edges[edge_id].state = NodeState.SKIPPED - - def mark_node_skipped(self, node_id: str) -> None: - self.graph.nodes[node_id].state = NodeState.SKIPPED - - def is_node_ready(self, node_id: str) -> bool: - return node_id == "selected" - - def analyze_edge_states(self, edges: list[Edge]) -> dict[str, bool]: - states = [edge.state for edge in edges] - return { - "has_unknown": any(state == NodeState.UNKNOWN for state in states), - "has_taken": any(state == NodeState.TAKEN for state in states), - "all_skipped": bool(states) - and all(state == NodeState.SKIPPED for state in states), - } - - def enqueue_node(self, node_id: str) -> None: - self.started_nodes.append(node_id) - - -def _branch_processor() -> EdgeProcessor: +def _branch_scheduler() -> Scheduler: graph = _Graph() - state_manager = _StateManager(graph) - skip_propagator = SkipPropagator( - graph=cast(Graph, graph), - state_manager=cast(GraphStateManager, state_manager), - ) - return EdgeProcessor( + return Scheduler( graph=cast(Graph, graph), - state_manager=cast(GraphStateManager, state_manager), - skip_propagator=skip_propagator, + state=RuntimeState( + workflow_id="workflow", variable_pool=VariablePool(), start_at=0 + ), + frame_id="root", ) @@ -121,10 +77,10 @@ def _edge_payloads( ] -def test_edge_processor_emits_taken_and_skipped_events_for_branch() -> None: - processor = _branch_processor() +def test_scheduler_emits_taken_and_skipped_events_for_branch() -> None: + scheduler = _branch_scheduler() - ready_nodes, events = processor.handle_branch_completion("branch", "yes") + ready_nodes, events = scheduler.handle_branch_completion("branch", "yes") assert ready_nodes == ["selected"] assert any(isinstance(event, GraphEdgeTakenEvent) for event in events) @@ -137,9 +93,9 @@ def test_edge_processor_emits_taken_and_skipped_events_for_branch() -> None: def test_process_node_success_emits_propagated_skip_events_for_branch() -> None: - processor = _branch_processor() + scheduler = _branch_scheduler() - ready_nodes, events = processor.process_node_success("branch", "yes") + ready_nodes, events = scheduler.process_node_success("branch", "yes") assert ready_nodes == ["selected"] assert _edge_payloads(events) == [ diff --git a/tests/graph_engine/test_serializable_graph_runtime.py b/tests/engine/test_serializable_graph_runtime.py similarity index 85% rename from tests/graph_engine/test_serializable_graph_runtime.py rename to tests/engine/test_serializable_graph_runtime.py index a426e524..32e897a9 100644 --- a/tests/graph_engine/test_serializable_graph_runtime.py +++ b/tests/engine/test_serializable_graph_runtime.py @@ -15,8 +15,29 @@ from graphon.dsl import inspect from graphon.dsl.entities import DslCredentials from graphon.dsl.node_factory import SlimDslNodeFactory +from graphon.engine import Engine +from graphon.engine.container_handler import LoopContainerHandler +from graphon.engine.frame import ExecutionFrame, FrameRegistry +from graphon.engine.ready_queue.entities import ResumeTask, StartTask +from graphon.engine.ready_queue.in_memory import InMemoryReadyQueue +from graphon.engine.scheduler import Scheduler +from graphon.engine.worker import DispatchTask, NodeEventTask, Worker +from graphon.engine_events.base import EngineEvent +from graphon.engine_events.graph import GraphRunPausedEvent, GraphRunStartedEvent +from graphon.engine_events.iteration import ( + NodeRunIterationStartedEvent, + NodeRunIterationSucceededEvent, +) +from graphon.engine_events.loop import ( + NodeRunLoopStartedEvent, + NodeRunLoopSucceededEvent, +) +from graphon.engine_events.node import ( + NodeRunPauseRequestedEvent, + NodeRunSucceededEvent, +) from graphon.entities.graph_config import NodeConfigDict -from graphon.entities.graph_init_params import GraphInitParams +from graphon.entities.graph_init_params import InitParams from graphon.entities.pause_reason import HitlRequired from graphon.entities.workflow_start_reason import WorkflowStartReason from graphon.enums import ( @@ -26,30 +47,6 @@ WorkflowNodeExecutionMetadataKey, ) from graphon.graph.graph import Graph -from graphon.graph_engine.command_channels.in_memory_channel import InMemoryChannel -from graphon.graph_engine.config import GraphEngineConfig -from graphon.graph_engine.entities.tasks import DispatchTask, TaskEvent -from graphon.graph_engine.frames import ExecutionFrame, FrameRegistry -from graphon.graph_engine.graph_engine import GraphEngine -from graphon.graph_engine.graph_state_manager import GraphStateManager -from graphon.graph_engine.loop_container_handler import LoopContainerHandler -from graphon.graph_engine.ready_queue.in_memory import InMemoryReadyQueue -from graphon.graph_engine.ready_queue.protocol import ResumeTask, StartTask -from graphon.graph_engine.worker import Worker -from graphon.graph_events.base import GraphEngineEvent -from graphon.graph_events.graph import GraphRunPausedEvent, GraphRunStartedEvent -from graphon.graph_events.iteration import ( - NodeRunIterationStartedEvent, - NodeRunIterationSucceededEvent, -) -from graphon.graph_events.loop import ( - NodeRunLoopStartedEvent, - NodeRunLoopSucceededEvent, -) -from graphon.graph_events.node import ( - NodeRunPauseRequestedEvent, - NodeRunSucceededEvent, -) from graphon.nodes.base.node import Node from graphon.nodes.container_effects import ( IterationFrameRequest, @@ -66,12 +63,13 @@ from graphon.nodes.loop.loop_node import LoopNode from graphon.runtime.container_state import ( FrameRuntimeData, + IterationFrameState, IterationRunState, LoopFrameState, LoopRunState, create_container_run_state, ) -from graphon.runtime.graph_runtime_state import GraphRuntimeState +from graphon.runtime.graph_runtime_state import RuntimeState from graphon.runtime.variable_pool import VariablePool from graphon.variables.segments import StringSegment from tests.helpers.workflow_events import final_outputs @@ -200,7 +198,7 @@ class _HitlNodeFactory: def with_runtime_state( self, - graph_runtime_state: GraphRuntimeState, + graph_runtime_state: RuntimeState, ) -> _HitlNodeFactory: return _HitlNodeFactory( base_factory=self.base_factory.with_runtime_state(graph_runtime_state), @@ -219,27 +217,27 @@ def create_node(self, node_config: NodeConfigDict) -> Node: ) -def _new_runtime_state(start_inputs: Mapping[str, object]) -> GraphRuntimeState: +def _new_runtime_state(start_inputs: Mapping[str, object]) -> RuntimeState: variable_pool = VariablePool() variable_pool.add(("sys", "workflow_execution_id"), "workflow-execution") for key, value in start_inputs.items(): variable_pool.add(("start", key), value) variable_pool.add(("sys", key), value) - return GraphRuntimeState(variable_pool=variable_pool, start_at=0) + return RuntimeState(workflow_id="workflow", variable_pool=variable_pool, start_at=0) def _hitl_engine( dsl: str, *, - runtime_state: GraphRuntimeState, + runtime_state: RuntimeState, callback: HITLCallback, -) -> GraphEngine: +) -> Engine: plan = inspect(dsl) graph_config = plan.document.graph_config if graph_config is None: msg = "test DSL must contain a graph" raise AssertionError(msg) - graph_init_params = GraphInitParams( + graph_init_params = InitParams( workflow_id="workflow", graph_config=graph_config, run_context={"workflow_execution_id": "workflow-execution"}, @@ -260,18 +258,16 @@ def _hitl_engine( ), root_node_id="start", ) - return GraphEngine( - workflow_id="workflow", + return Engine( graph=graph, graph_runtime_state=runtime_state, - command_channel=InMemoryChannel(), - config=GraphEngineConfig(min_workers=2, max_workers=2), + workers=2, ) def _snapshot_after_hitl_pause( - engine: GraphEngine, -) -> tuple[str, list[GraphEngineEvent]]: + engine: Engine, +) -> tuple[str, list[EngineEvent]]: events = list(engine.run()) assert any(isinstance(event, NodeRunPauseRequestedEvent) for event in events) assert any(isinstance(event, GraphRunPausedEvent) for event in events) @@ -301,22 +297,21 @@ def _execution_frame( *, frame_id: str, graph: Graph, - graph_runtime_state: GraphRuntimeState, + graph_runtime_state: RuntimeState, ) -> ExecutionFrame: return ExecutionFrame( frame_id=frame_id, graph=graph, - graph_runtime_state=graph_runtime_state, - state_manager=GraphStateManager(graph, graph_runtime_state, frame_id), - edge_processor=cast(Any, SimpleNamespace()), - error_handler=cast(Any, SimpleNamespace()), + state=graph_runtime_state, + scheduler=Scheduler(graph, graph_runtime_state, frame_id), + failure_handler=cast(Any, SimpleNamespace()), ) class _FrameFactory: def with_runtime_state( self, - graph_runtime_state: GraphRuntimeState, + graph_runtime_state: RuntimeState, ) -> _FrameFactory: _ = graph_runtime_state return self @@ -340,10 +335,16 @@ class _GraphNode: state = NodeState.UNKNOWN -def _loop_graph(runtime_state: GraphRuntimeState) -> Graph: +def _loop_graph(runtime_state: RuntimeState) -> Graph: graph_config = { "nodes": [ - {"id": "loop-start", "data": {"type": BuiltinNodeTypes.LOOP_START}}, + { + "id": "loop-start", + "data": { + "type": BuiltinNodeTypes.LOOP_START, + "container_id": "loop", + }, + }, ], "edges": [], } @@ -370,9 +371,10 @@ def _loop_graph(runtime_state: GraphRuntimeState) -> Graph: ) -def _runtime_with_live_resume_task() -> GraphRuntimeState: +def _runtime_with_live_resume_task() -> RuntimeState: ready_queue = InMemoryReadyQueue() - graph_runtime_state = GraphRuntimeState( + graph_runtime_state = RuntimeState( + workflow_id="workflow", variable_pool=VariablePool(), start_at=1, ready_queue=ready_queue, @@ -407,7 +409,7 @@ def _runtime_with_live_resume_task() -> GraphRuntimeState: loop_handler = LoopContainerHandler( frame_registry=frame_registry, ) - loop_handler.start_await( + loop_handler.handle_request( invocation_id="loop-invocation", request=request, ) @@ -415,18 +417,18 @@ def _runtime_with_live_resume_task() -> GraphRuntimeState: frame_id="loop-invocation:loop:0", node_id="loop-start", ) - child_frame = frame_registry.get("loop-invocation:loop:0") - child_frame.state_manager.finish_execution("loop-start") + child_frame = frame_registry["loop-invocation:loop:0"] + child_frame.scheduler.finish_execution("loop-start") - loop_handler.complete_frame(child_frame) + loop_handler.complete_frame_if_ready(child_frame) resume_task = ready_queue.get(timeout=0.01) assert isinstance(resume_task, ResumeTask) ready_queue.put(resume_task) return graph_runtime_state -def _resume_loop_snapshot(snapshot: str) -> list[TaskEvent]: - runtime_state = GraphRuntimeState.from_snapshot(snapshot) +def _resume_loop_snapshot(snapshot: str) -> list[NodeEventTask]: + runtime_state = RuntimeState.from_snapshot(snapshot) runtime_state.graph_execution.paused = False for task in runtime_state.drain_deferred_ready_tasks(): runtime_state.enqueue_ready_task(task) @@ -439,12 +441,12 @@ def _resume_loop_snapshot(snapshot: str) -> list[TaskEvent]: graph_runtime_state=runtime_state, ), ) - event_queue: queue.Queue[DispatchTask] = queue.Queue() + dispatch_queue: queue.Queue[DispatchTask] = queue.Queue() task_claiming = Event() task_claiming.set() worker = Worker( ready_queue=cast(InMemoryReadyQueue, runtime_state.ready_queue), - event_queue=event_queue, + dispatch_queue=dispatch_queue, frame_registry=frame_registry, layers=[], task_claim_lock=Lock(), @@ -452,10 +454,10 @@ def _resume_loop_snapshot(snapshot: str) -> list[TaskEvent]: ) worker.start() try: - first_event = event_queue.get(timeout=1) - second_event = event_queue.get(timeout=1) - assert isinstance(first_event, TaskEvent) - assert isinstance(second_event, TaskEvent) + first_event = dispatch_queue.get(timeout=1) + second_event = dispatch_queue.get(timeout=1) + assert isinstance(first_event, NodeEventTask) + assert isinstance(second_event, NodeEventTask) return [first_event, second_event] finally: worker.stop() @@ -463,7 +465,8 @@ def _resume_loop_snapshot(snapshot: str) -> list[TaskEvent]: def test_resume_restores_container_runs_before_workers_start() -> None: - runtime_state = GraphRuntimeState( + runtime_state = RuntimeState( + workflow_id="workflow", variable_pool=VariablePool(), start_at=1, ) @@ -486,20 +489,20 @@ def test_resume_restores_container_runs_before_workers_start() -> None: ), ) runtime_state.defer_ready_task(StartTask(frame_id="root", node_id="start")) - state_manager = MagicMock() + scheduler = MagicMock() worker_pool = MagicMock() def assert_tasks_are_tracked_before_workers_start() -> None: assert runtime_state.ready_queue.qsize() == 0 - assert state_manager.track_unfinished.call_args_list == [ + assert scheduler.track_unfinished.call_args_list == [ call("loop"), call("start"), ] worker_pool.start.side_effect = assert_tasks_are_tracked_before_workers_start frame_registry = MagicMock() - frame_registry.get.return_value.state_manager = state_manager - engine = object.__new__(GraphEngine) + frame_registry.__getitem__.return_value.scheduler = scheduler + engine = object.__new__(Engine) engine._graph_runtime_state = runtime_state engine._frame_registry = frame_registry engine._worker_pool = worker_pool @@ -512,10 +515,12 @@ def assert_tasks_are_tracked_before_workers_start() -> None: engine._dispatcher.start.assert_called_once_with() -def test_loop_frame_restore_shares_parent_variable_pool() -> None: +def test_loop_frame_restore_copies_parent_variable_pool() -> None: parent_pool = VariablePool() parent_pool.add(["loop", "seed"], "parent") - runtime_state = GraphRuntimeState(variable_pool=parent_pool, start_at=1) + runtime_state = RuntimeState( + workflow_id="workflow", variable_pool=parent_pool, start_at=1 + ) request = LoopFrameRequest( inputs={"loop_count": build_container_value(1)}, outputs={}, @@ -551,7 +556,7 @@ def test_loop_frame_restore_shares_parent_variable_pool() -> None: ) runtime_state.put_container_frame(frame_state) - restored_state = GraphRuntimeState.from_snapshot(runtime_state.dumps()) + restored_state = RuntimeState.from_snapshot(runtime_state.dumps()) restored_frame_state = restored_state.get_container_frame(frame_state.frame_id) frame_registry = FrameRegistry() frame_registry.register( @@ -566,14 +571,17 @@ def test_loop_frame_restore_shares_parent_variable_pool() -> None: restored_frame_state, ) - restored_pool = frame_registry.get( - restored_frame_state.frame_id, - ).graph_runtime_state.variable_pool - assert restored_frame_state.runtime_data.variable_pool == "parent" - assert restored_pool is restored_state.variable_pool + restored_pool = frame_registry[restored_frame_state.frame_id].state.variable_pool + restored_frame_state = restored_state.get_container_frame(frame_state.frame_id) + assert not isinstance(restored_frame_state.runtime_data.variable_pool, str) + assert restored_pool is not restored_state.variable_pool restored_seed = restored_pool.get(["loop", "seed"]) assert restored_seed is not None assert restored_seed.to_object() == "parent" + restored_pool.add(["loop", "seed"], "child") + parent_seed = restored_state.variable_pool.get(["loop", "seed"]) + assert parent_seed is not None + assert parent_seed.to_object() == "parent" def test_loop_hitl_runtime_state_round_trip_preserves_progress() -> None: @@ -593,14 +601,14 @@ def pause_after_one_round(context: HITLContext) -> Completed | PauseRequested: callback=pause_after_one_round, ) ) - paused_state = GraphRuntimeState.from_snapshot(snapshot) + paused_state = RuntimeState.from_snapshot(snapshot) run_state = paused_state.container_runs()[0] frame_state = paused_state.container_frames()[0] deferred_tasks = paused_state.drain_deferred_ready_tasks() resumed_events = list( _hitl_engine( _loop_dsl(), - runtime_state=GraphRuntimeState.from_snapshot(snapshot), + runtime_state=RuntimeState.from_snapshot(snapshot), callback=_complete_loop_hitl, ).run() ) @@ -610,7 +618,7 @@ def pause_after_one_round(context: HITLContext) -> Completed | PauseRequested: if isinstance(event, NodeRunPauseRequestedEvent) ] assert len(pause_requests) == 1 - assert pause_requests[0].in_loop_id == "loop" + assert pause_requests[0].container_id == "loop" assert pause_requests[0].reason == HitlRequired( session_id="session-human-input", node_id="human-input", @@ -621,14 +629,14 @@ def pause_after_one_round(context: HITLContext) -> Completed | PauseRequested: for event in paused_events if isinstance(event, NodeRunSucceededEvent) and event.node_id == "human-input" - and event.in_loop_id == "loop" + and event.container_id == "loop" ] resumed_successes = [ event for event in resumed_events if isinstance(event, NodeRunSucceededEvent) and event.node_id == "human-input" - and event.in_loop_id == "loop" + and event.container_id == "loop" ] loop_started = next( event for event in paused_events if isinstance(event, NodeRunLoopStartedEvent) @@ -645,6 +653,7 @@ def pause_after_one_round(context: HITLContext) -> Completed | PauseRequested: "seed": "fixed", "loop_round": 1, } + assert isinstance(frame_state, LoopFrameState) assert frame_state.index == 1 assert deferred_tasks == [ StartTask(frame_id=frame_state.frame_id, node_id="human-input") @@ -721,14 +730,14 @@ def pause_with_active_sibling( callback=pause_with_active_sibling, ) ) - paused_state = GraphRuntimeState.from_snapshot(snapshot) + paused_state = RuntimeState.from_snapshot(snapshot) run_state = paused_state.container_runs()[0] frame_state = paused_state.container_frames()[0] deferred_tasks = paused_state.drain_deferred_ready_tasks() resumed_events = list( _hitl_engine( _iteration_dsl(), - runtime_state=GraphRuntimeState.from_snapshot(snapshot), + runtime_state=RuntimeState.from_snapshot(snapshot), callback=_complete_iteration_hitl, ).run() ) @@ -742,19 +751,19 @@ def pause_with_active_sibling( for event in paused_events if isinstance(event, NodeRunSucceededEvent) and event.node_id == "human-input" - and event.in_iteration_id == "iteration" + and event.container_id == "iteration" ] resumed_successes = [ event for event in resumed_events if isinstance(event, NodeRunSucceededEvent) and event.node_id == "human-input" - and event.in_iteration_id == "iteration" + and event.container_id == "iteration" ] start_tasks = [task for task in deferred_tasks if isinstance(task, StartTask)] resume_tasks = [task for task in deferred_tasks if isinstance(task, ResumeTask)] assert len(pause_requests) == 1 - assert pause_requests[0].in_iteration_id == "iteration" + assert pause_requests[0].container_id == "iteration" assert pause_requests[0].reason == HitlRequired( session_id="session-beta", node_id="human-input", @@ -769,6 +778,7 @@ def pause_with_active_sibling( assert {key: value.to_object() for key, value in run_state.outputs.items()} == { "0": "alpha!" } + assert isinstance(frame_state, IterationFrameState) assert frame_state.index == 1 assert len(start_tasks) == 1 assert start_tasks[0] == StartTask( @@ -807,7 +817,7 @@ def test_deferred_resume_task_round_trips_and_resumes_parent_container() -> None for task in runtime_state.ready_queue.drain(): runtime_state.defer_ready_task(task) snapshot = runtime_state.dumps() - restored_for_assert = GraphRuntimeState.from_snapshot(snapshot) + restored_for_assert = RuntimeState.from_snapshot(snapshot) deferred_tasks = restored_for_assert.drain_deferred_ready_tasks() task_events = _resume_loop_snapshot(snapshot) diff --git a/tests/graph_events/__init__.py b/tests/engine_events/__init__.py similarity index 100% rename from tests/graph_events/__init__.py rename to tests/engine_events/__init__.py diff --git a/tests/graph_events/test_traversal_events.py b/tests/engine_events/test_traversal_events.py similarity index 85% rename from tests/graph_events/test_traversal_events.py rename to tests/engine_events/test_traversal_events.py index e0044610..a0ac9e57 100644 --- a/tests/graph_events/test_traversal_events.py +++ b/tests/engine_events/test_traversal_events.py @@ -1,4 +1,4 @@ -from graphon.graph_events import GraphEdgeSkippedEvent, GraphEdgeTakenEvent +from graphon.engine_events import GraphEdgeSkippedEvent, GraphEdgeTakenEvent def test_graph_edge_taken_event_exports_payload() -> None: @@ -14,6 +14,7 @@ def test_graph_edge_taken_event_exports_payload() -> None: "source_node_id": "source", "target_node_id": "target", "source_handle": "success", + "container_id": "", } @@ -29,4 +30,5 @@ def test_graph_edge_skipped_event_exports_payload() -> None: "source_node_id": "source", "target_node_id": "other", "source_handle": None, + "container_id": "", } diff --git a/tests/graph/test_graph_scoping.py b/tests/graph/test_graph_scoping.py new file mode 100644 index 00000000..e5a1efbe --- /dev/null +++ b/tests/graph/test_graph_scoping.py @@ -0,0 +1,245 @@ +from __future__ import annotations + +from collections.abc import Mapping +from dataclasses import dataclass, field +from typing import Any, cast +from unittest.mock import Mock + +import pytest + +from graphon.entities.graph_config import NodeConfigDict +from graphon.enums import NodeExecutionType, NodeState +from graphon.graph.graph import Graph +from graphon.graph.validation import GraphValidationError +from graphon.nodes.base.node import Node + + +@dataclass(slots=True) +class _RecordingNodeFactory: + created_node_ids: list[str] = field(default_factory=list) + + def create_node(self, node_config: NodeConfigDict) -> Node: + node_id = node_config["id"] + self.created_node_ids.append(node_id) + + data = node_config["data"] + node_type = data.get("type") + node = cast(Any, Mock(spec=Node)) + node.id = node_id + node.node_type = node_type + node.execution_type = ( + NodeExecutionType.ROOT + if node_type in {"start", "iteration-start", "loop-start"} + else NodeExecutionType.EXECUTABLE + ) + node.error_strategy = None + node.state = NodeState.UNKNOWN + return node + + +def _node( + node_id: str, + *, + node_type: str = "answer", + container_id: str = "", +) -> dict[str, Any]: + return { + "id": node_id, + "data": {"type": node_type, "container_id": container_id}, + } + + +def _edge(source: str, target: str) -> dict[str, str]: + return {"source": source, "target": target} + + +def _scoped_graph_config() -> dict[str, Any]: + return { + "nodes": [ + _node("start", node_type="start"), + _node("container-a", node_type="loop"), + _node("container-b", node_type="iteration"), + _node("end"), + _node("a-start", node_type="start", container_id="container-a"), + _node( + "nested-container", + node_type="iteration", + container_id="container-a", + ), + _node("a-end", container_id="container-a"), + _node( + "nested-start", + node_type="start", + container_id="nested-container", + ), + _node("nested-end", container_id="nested-container"), + _node("b-start", node_type="start", container_id="container-b"), + _node("b-end", container_id="container-b"), + ], + "edges": [ + _edge("start", "container-a"), + _edge("container-a", "container-b"), + _edge("container-b", "end"), + _edge("a-start", "nested-container"), + _edge("nested-container", "a-end"), + _edge("nested-start", "nested-end"), + _edge("b-start", "b-end"), + ], + "viewport": {"x": 10, "y": 20}, + } + + +def _config_node_ids(graph_config: Mapping[str, Any]) -> set[str]: + return {node["id"] for node in graph_config["nodes"]} + + +def _config_edges(graph_config: Mapping[str, Any]) -> set[tuple[str, str]]: + return {(edge["source"], edge["target"]) for edge in graph_config["edges"]} + + +def _graph_edges(graph: Graph) -> set[tuple[str, str]]: + return {(edge.tail, edge.head) for edge in graph.edges.values()} + + +def test_graph_init_materializes_only_root_scope_by_default() -> None: + graph_config = _scoped_graph_config() + node_factory = _RecordingNodeFactory() + + graph = Graph.init( + graph_config=graph_config, + node_factory=node_factory, + root_node_id="start", + ) + + assert set(graph.nodes) == {"start", "container-a", "container-b", "end"} + assert set(node_factory.created_node_ids) == set(graph.nodes) + assert _graph_edges(graph) == { + ("start", "container-a"), + ("container-a", "container-b"), + ("container-b", "end"), + } + assert graph.graph_config is not None + assert all(node.graph_config is graph.graph_config for node in graph.nodes.values()) + assert _config_node_ids(graph.graph_config) == { + node["id"] for node in graph_config["nodes"] + } + assert graph.graph_config["viewport"] == graph_config["viewport"] + + +def test_graph_init_scopes_execution_and_retains_only_its_subtree_config() -> None: + graph_config = _scoped_graph_config() + node_factory = _RecordingNodeFactory() + + graph = Graph.init( + graph_config=graph_config, + node_factory=node_factory, + root_node_id="a-start", + container_id="container-a", + ) + + assert set(graph.nodes) == {"a-start", "nested-container", "a-end"} + assert _graph_edges(graph) == { + ("a-start", "nested-container"), + ("nested-container", "a-end"), + } + assert graph.graph_config is not None + assert all(node.graph_config is graph.graph_config for node in graph.nodes.values()) + assert _config_node_ids(graph.graph_config) == { + "a-start", + "nested-container", + "a-end", + "nested-start", + "nested-end", + } + assert _config_edges(graph.graph_config) == { + ("a-start", "nested-container"), + ("nested-container", "a-end"), + ("nested-start", "nested-end"), + } + + +def test_graph_init_rejects_edges_crossing_container_scopes() -> None: + graph_config = _scoped_graph_config() + graph_config["edges"].append(_edge("container-a", "a-start")) + + with pytest.raises( + ValueError, + match=( + r"Edge 'container-a->a-start' crosses container scopes " + r"'' and 'container-a'" + ), + ): + Graph.init( + graph_config=graph_config, + node_factory=_RecordingNodeFactory(), + root_node_id="start", + ) + + +def test_graph_init_rejects_orphan_scopes_and_unknown_edges() -> None: + graph_config = _scoped_graph_config() + graph_config["nodes"].append(_node("orphan", container_id="missing")) + with pytest.raises(ValueError, match="orphan"): + Graph.init( + graph_config=graph_config, + node_factory=_RecordingNodeFactory(), + root_node_id="start", + ) + + graph_config = _scoped_graph_config() + graph_config["edges"].append(_edge("ghost-a", "ghost-b")) + with pytest.raises(GraphValidationError): + Graph.init( + graph_config=graph_config, + node_factory=_RecordingNodeFactory(), + root_node_id="start", + ) + + +def test_graph_init_requires_container_id_for_nested_legacy_fields() -> None: + graph_config = { + "nodes": [ + {"id": "start", "data": {"type": "start"}}, + {"id": "loop", "data": {"type": "loop"}}, + { + "id": "loop-start", + "data": {"type": "loop-start", "loop_id": "loop"}, + }, + { + "id": "iteration", + "data": {"type": "iteration", "loop_id": "loop"}, + }, + { + "id": "iteration-start", + "data": { + "type": "iteration-start", + "loop_id": "loop", + "iteration_id": "iteration", + }, + }, + { + "id": "stop", + "data": { + "type": "loop-end", + "loop_id": "loop", + "iteration_id": "iteration", + }, + }, + ], + "edges": [ + _edge("start", "loop"), + _edge("loop-start", "iteration"), + _edge("iteration-start", "stop"), + ], + } + + with pytest.raises( + ValueError, + match=r"Nested container nodes must set data\.container_id", + ): + Graph.init( + graph_config=graph_config, + node_factory=_RecordingNodeFactory(), + root_node_id="iteration-start", + container_id="iteration", + ) diff --git a/tests/graph/test_graph_validation.py b/tests/graph/test_graph_validation.py index 910e5fe3..6d340539 100644 --- a/tests/graph/test_graph_validation.py +++ b/tests/graph/test_graph_validation.py @@ -8,13 +8,13 @@ from graphon.entities.base_node_data import BaseNodeData from graphon.entities.graph_config import NodeConfigDict -from graphon.entities.graph_init_params import GraphInitParams +from graphon.entities.graph_init_params import InitParams from graphon.enums import BuiltinNodeTypes, ErrorStrategy, NodeExecutionType, NodeType from graphon.graph.graph import Graph from graphon.graph.validation import GraphValidationError -from graphon.node_events.base import NodeEventBase, NodeRunResult +from graphon.node_events.base import NodeEventPayload, NodeRunResult from graphon.nodes.base.node import Node -from graphon.runtime.graph_runtime_state import GraphRuntimeState +from graphon.runtime.graph_runtime_state import RuntimeState from graphon.runtime.variable_pool import VariablePool from ..helpers import build_graph_init_params @@ -38,8 +38,8 @@ def __init__( *, node_id: str, data: _TestNodeData, - graph_init_params: GraphInitParams, - graph_runtime_state: GraphRuntimeState, + graph_init_params: InitParams, + graph_runtime_state: RuntimeState, ) -> None: super().__init__( node_id=node_id, @@ -52,7 +52,7 @@ def __init__( if isinstance(node_type_value, str): self.node_type = node_type_value - def _run(self) -> NodeRunResult | Generator[NodeEventBase, None, None]: + def _run(self) -> NodeRunResult | Generator[NodeEventPayload, None, None]: raise NotImplementedError def post_init(self) -> None: @@ -72,8 +72,8 @@ def _maybe_override_execution_type(self) -> None: @dataclass(slots=True) class _SimpleNodeFactory: - graph_init_params: GraphInitParams - graph_runtime_state: GraphRuntimeState + graph_init_params: InitParams + graph_runtime_state: RuntimeState def create_node(self, node_config: NodeConfigDict) -> _TestNode: return _TestNode( @@ -91,7 +91,8 @@ def graph_init_dependencies() -> tuple[_SimpleNodeFactory, dict[str, object]]: workflow_id="workflow", graph_config=graph_config, ) - runtime_state = GraphRuntimeState( + runtime_state = RuntimeState( + workflow_id="workflow", variable_pool=VariablePool(), start_at=time.perf_counter(), ) diff --git a/tests/graph_engine/__init__.py b/tests/graph_engine/__init__.py deleted file mode 100644 index e69de29b..00000000 diff --git a/tests/graph_engine/graph_traversal/__init__.py b/tests/graph_engine/graph_traversal/__init__.py deleted file mode 100644 index e69de29b..00000000 diff --git a/tests/graph_engine/graph_traversal/test_skip_propagator.py b/tests/graph_engine/graph_traversal/test_skip_propagator.py deleted file mode 100644 index 12c73e37..00000000 --- a/tests/graph_engine/graph_traversal/test_skip_propagator.py +++ /dev/null @@ -1,225 +0,0 @@ -from unittest.mock import MagicMock, create_autospec - -from graphon.graph.edge import Edge -from graphon.graph.graph import Graph -from graphon.graph_engine.graph_state_manager import GraphStateManager -from graphon.graph_engine.graph_traversal.skip_propagator import SkipPropagator - - -class TestSkipPropagator: - def test_propagate_skip_from_edge_with_unknown_edges_stops_processing(self) -> None: - mock_graph = create_autospec(Graph) - mock_state_manager = create_autospec(GraphStateManager) - - mock_edge = MagicMock(spec=Edge) - mock_edge.id = "edge_1" - mock_edge.head = "node_2" - - mock_graph.edges = {"edge_1": mock_edge} - - incoming_edges = [MagicMock(spec=Edge), MagicMock(spec=Edge)] - mock_graph.get_incoming_edges.return_value = incoming_edges - - mock_state_manager.analyze_edge_states.return_value = { - "has_unknown": True, - "has_taken": False, - "all_skipped": False, - } - - propagator = SkipPropagator(mock_graph, mock_state_manager) - - propagator.propagate_skip_from_edge("edge_1") - - mock_graph.get_incoming_edges.assert_called_once_with("node_2") - mock_state_manager.analyze_edge_states.assert_called_once_with(incoming_edges) - mock_state_manager.enqueue_node.assert_not_called() - mock_state_manager.mark_node_skipped.assert_not_called() - - def test_propagate_skip_from_edge_with_taken_edge_enqueues_node(self) -> None: - mock_graph = create_autospec(Graph) - mock_state_manager = create_autospec(GraphStateManager) - - mock_edge = MagicMock(spec=Edge) - mock_edge.id = "edge_1" - mock_edge.head = "node_2" - - mock_graph.edges = {"edge_1": mock_edge} - incoming_edges = [MagicMock(spec=Edge)] - mock_graph.get_incoming_edges.return_value = incoming_edges - - mock_state_manager.analyze_edge_states.return_value = { - "has_unknown": False, - "has_taken": True, - "all_skipped": False, - } - - propagator = SkipPropagator(mock_graph, mock_state_manager) - - propagator.propagate_skip_from_edge("edge_1") - - mock_state_manager.enqueue_node.assert_called_once_with("node_2") - mock_state_manager.mark_node_skipped.assert_not_called() - - def test_propagate_skip_from_edge_with_all_skipped_propagates_to_node(self) -> None: - mock_graph = create_autospec(Graph) - mock_state_manager = create_autospec(GraphStateManager) - - mock_edge = MagicMock(spec=Edge) - mock_edge.id = "edge_1" - mock_edge.head = "node_2" - - mock_graph.edges = {"edge_1": mock_edge} - incoming_edges = [MagicMock(spec=Edge)] - mock_graph.get_incoming_edges.return_value = incoming_edges - - mock_state_manager.analyze_edge_states.return_value = { - "has_unknown": False, - "has_taken": False, - "all_skipped": True, - } - - propagator = SkipPropagator(mock_graph, mock_state_manager) - - propagator.propagate_skip_from_edge("edge_1") - - mock_state_manager.mark_node_skipped.assert_called_once_with("node_2") - mock_state_manager.enqueue_node.assert_not_called() - - def test_propagate_skip_to_node_marks_node_and_outgoing_edges_skipped(self) -> None: - mock_graph = create_autospec(Graph) - mock_state_manager = create_autospec(GraphStateManager) - - edge1 = MagicMock(spec=Edge) - edge1.id = "edge_2" - edge1.tail = "node_1" - edge1.head = "node_downstream_1" - edge1.source_handle = "source" - - edge2 = MagicMock(spec=Edge) - edge2.id = "edge_3" - edge2.tail = "node_1" - edge2.head = "node_downstream_2" - edge2.source_handle = "source" - - mock_graph.edges = {"edge_2": edge1, "edge_3": edge2} - mock_graph.get_outgoing_edges.return_value = [edge1, edge2] - mock_graph.get_incoming_edges.return_value = [] - - propagator = SkipPropagator(mock_graph, mock_state_manager) - - propagator.propagate_skip_to_node("node_1") - - mock_state_manager.mark_node_skipped.assert_called_once_with("node_1") - mock_state_manager.mark_edge_skipped.assert_any_call("edge_2") - mock_state_manager.mark_edge_skipped.assert_any_call("edge_3") - assert mock_state_manager.mark_edge_skipped.call_count == 2 - - def test_skip_branch_paths_marks_unselected_edges_and_propagates(self) -> None: - mock_graph = create_autospec(Graph) - mock_state_manager = create_autospec(GraphStateManager) - - edge1 = MagicMock(spec=Edge) - edge1.id = "edge_1" - edge1.tail = "node_1" - edge1.head = "node_downstream_1" - edge1.source_handle = "source" - - edge2 = MagicMock(spec=Edge) - edge2.id = "edge_2" - edge2.tail = "node_1" - edge2.head = "node_downstream_2" - edge2.source_handle = "source" - - mock_graph.edges = {"edge_1": edge1, "edge_2": edge2} - mock_graph.get_incoming_edges.return_value = [] - - propagator = SkipPropagator(mock_graph, mock_state_manager) - - propagator.skip_branch_paths([edge1, edge2]) - - mock_state_manager.mark_edge_skipped.assert_any_call("edge_1") - mock_state_manager.mark_edge_skipped.assert_any_call("edge_2") - assert mock_state_manager.mark_edge_skipped.call_count == 2 - - def test_propagate_skip_from_edge_recursively_propagates_through_graph( - self, - ) -> None: - mock_graph = create_autospec(Graph) - mock_state_manager = create_autospec(GraphStateManager) - - edge1 = MagicMock(spec=Edge) - edge1.id = "edge_1" - edge1.tail = "node_1" - edge1.head = "node_2" - edge1.source_handle = "source" - - edge3 = MagicMock(spec=Edge) - edge3.id = "edge_3" - edge3.tail = "node_2" - edge3.head = "node_4" - edge3.source_handle = "source" - - mock_graph.edges = {"edge_1": edge1, "edge_3": edge3} - - def get_incoming_edges_side_effect(node_id: str) -> list[Edge]: - if node_id == "node_2": - return [edge1] - if node_id == "node_4": - return [edge3] - return [] - - mock_graph.get_incoming_edges.side_effect = get_incoming_edges_side_effect - - def get_outgoing_edges_side_effect(node_id: str) -> list[Edge]: - if node_id == "node_2": - return [edge3] - if node_id == "node_4": - return [] - return [] - - mock_graph.get_outgoing_edges.side_effect = get_outgoing_edges_side_effect - - mock_state_manager.analyze_edge_states.return_value = { - "has_unknown": False, - "has_taken": False, - "all_skipped": True, - } - - propagator = SkipPropagator(mock_graph, mock_state_manager) - - propagator.propagate_skip_from_edge("edge_1") - - mock_state_manager.mark_node_skipped.assert_any_call("node_2") - mock_state_manager.mark_edge_skipped.assert_any_call("edge_3") - mock_state_manager.mark_node_skipped.assert_any_call("node_4") - assert mock_state_manager.mark_node_skipped.call_count == 2 - - def test_propagate_skip_from_edge_with_mixed_edge_states_handles_correctly( - self, - ) -> None: - mock_graph = create_autospec(Graph) - mock_state_manager = create_autospec(GraphStateManager) - - mock_edge = MagicMock(spec=Edge) - mock_edge.id = "edge_1" - mock_edge.head = "node_2" - - mock_graph.edges = {"edge_1": mock_edge} - incoming_edges = [ - MagicMock(spec=Edge), - MagicMock(spec=Edge), - MagicMock(spec=Edge), - ] - mock_graph.get_incoming_edges.return_value = incoming_edges - - mock_state_manager.analyze_edge_states.return_value = { - "has_unknown": True, - "has_taken": False, - "all_skipped": False, - } - - propagator = SkipPropagator(mock_graph, mock_state_manager) - - propagator.propagate_skip_from_edge("edge_1") - mock_state_manager.enqueue_node.assert_not_called() - mock_state_manager.mark_node_skipped.assert_not_called() diff --git a/tests/helpers/builders.py b/tests/helpers/builders.py index 37d5b4f4..01c6a2b0 100644 --- a/tests/helpers/builders.py +++ b/tests/helpers/builders.py @@ -5,7 +5,7 @@ from collections.abc import Mapping, Sequence from typing import Any -from graphon.entities.graph_init_params import GraphInitParams +from graphon.entities.graph_init_params import InitParams from graphon.runtime.variable_pool import VariablePool from graphon.variables.variables import Variable @@ -29,8 +29,8 @@ def build_graph_init_params( graph_config: Mapping[str, Any] | None = None, run_context: Mapping[str, Any] | None = None, call_depth: int = 0, -) -> GraphInitParams: - return GraphInitParams( +) -> InitParams: + return InitParams( workflow_id=workflow_id, graph_config=graph_config or {}, run_context=run_context or {}, diff --git a/tests/helpers/workflow_events.py b/tests/helpers/workflow_events.py index 8b32d828..6cd36f51 100644 --- a/tests/helpers/workflow_events.py +++ b/tests/helpers/workflow_events.py @@ -8,9 +8,9 @@ import graphon.dsl.node_factory as node_factory_module from graphon.dsl import loads -from graphon.graph_events.base import GraphEngineEvent, GraphNodeEventBase -from graphon.graph_events.graph import GraphRunSucceededEvent -from graphon.graph_events.traversal import GraphEdgeTakenEvent +from graphon.engine_events.base import EngineEvent, NodeEvent +from graphon.engine_events.graph import GraphRunSucceededEvent +from graphon.engine_events.traversal import GraphEdgeTakenEvent from graphon.model_runtime.entities.common_entities import I18nObject from graphon.model_runtime.entities.llm_entities import LLMResult, LLMUsage from graphon.model_runtime.entities.message_entities import ( @@ -139,7 +139,7 @@ def run_workflow( *, start_inputs: Mapping[str, Any] = _EMPTY_MAPPING, credentials: Mapping[str, Any] = _EMPTY_MAPPING, -) -> list[GraphEngineEvent]: +) -> list[EngineEvent]: engine = loads( dsl, start_inputs=dict(start_inputs), @@ -149,31 +149,29 @@ def run_workflow( def event_path( - events: Sequence[GraphEngineEvent], -) -> list[tuple[str, str, str, str]]: + events: Sequence[EngineEvent], +) -> list[tuple[str, str, str]]: """Project events to stable fields while preserving the complete event order.""" - path: list[tuple[str, str, str, str]] = [] + path: list[tuple[str, str, str]] = [] for event in events: - if isinstance(event, GraphNodeEventBase): + if isinstance(event, NodeEvent): path.append(( type(event).__name__, event.node_id, - event.in_loop_id or "", - event.in_iteration_id or "", + event.container_id, )) elif isinstance(event, GraphEdgeTakenEvent): path.append(( type(event).__name__, f"{event.source_node_id}->{event.target_node_id}", - "", - "", + event.container_id, )) else: - path.append((type(event).__name__, "", "", "")) + path.append((type(event).__name__, "", "")) return path -def final_outputs(events: Sequence[GraphEngineEvent]) -> dict[str, object]: +def final_outputs(events: Sequence[EngineEvent]) -> dict[str, object]: for event in reversed(events): if isinstance(event, GraphRunSucceededEvent): return event.outputs diff --git a/tests/http/test_client.py b/tests/http/test_client.py index 083d9174..f518a42c 100644 --- a/tests/http/test_client.py +++ b/tests/http/test_client.py @@ -36,7 +36,7 @@ from graphon.nodes.question_classifier.question_classifier_node import ( QuestionClassifierNode, ) -from graphon.runtime.graph_runtime_state import GraphRuntimeState +from graphon.runtime.graph_runtime_state import RuntimeState from ..helpers import build_graph_init_params, build_variable_pool @@ -111,8 +111,9 @@ def _raise(self, method: str, url: str, **kwargs: Any) -> HttpResponse: raise AssertionError(msg) -def _build_runtime_state() -> GraphRuntimeState: - return GraphRuntimeState( +def _build_runtime_state() -> RuntimeState: + return RuntimeState( + workflow_id="workflow", variable_pool=build_variable_pool(), start_at=time.perf_counter(), ) diff --git a/tests/node_events/test_node_event_aliases.py b/tests/node_events/test_node_event_aliases.py index 97f2c02f..8a0e30b3 100644 --- a/tests/node_events/test_node_event_aliases.py +++ b/tests/node_events/test_node_event_aliases.py @@ -1,9 +1,9 @@ -from graphon.entities.pause_reason import SchedulingPause -from graphon.enums import WorkflowNodeExecutionStatus -from graphon.graph_events.node import ( +from graphon.engine_events.node import ( NodeRunPauseRequestedEvent, NodeRunVariableUpdatedEvent, ) +from graphon.entities.pause_reason import SchedulingPause +from graphon.enums import WorkflowNodeExecutionStatus from graphon.node_events.base import NodeRunResult from graphon.node_events.node import PauseRequestedEvent, VariableUpdatedEvent diff --git a/tests/nodes/base/test_node_execution_binding.py b/tests/nodes/base/test_node_execution_binding.py index 146a6054..ef210631 100644 --- a/tests/nodes/base/test_node_execution_binding.py +++ b/tests/nodes/base/test_node_execution_binding.py @@ -4,7 +4,7 @@ from graphon.enums import BuiltinNodeTypes, WorkflowNodeExecutionStatus from graphon.node_events.base import NodeRunResult from graphon.nodes.base.node import Node -from graphon.runtime.graph_runtime_state import GraphRuntimeState +from graphon.runtime.graph_runtime_state import RuntimeState from tests.helpers import build_graph_init_params, build_variable_pool @@ -24,7 +24,8 @@ def test_node_run_requires_bound_execution_id() -> None: node_id="node", data=BaseNodeData(type=BuiltinNodeTypes.CODE), graph_init_params=build_graph_init_params(), - graph_runtime_state=GraphRuntimeState( + graph_runtime_state=RuntimeState( + workflow_id="workflow", variable_pool=build_variable_pool(), start_at=1, ), diff --git a/tests/nodes/human_input/test_human_input_node.py b/tests/nodes/human_input/test_human_input_node.py index 519170fe..52e9cac1 100644 --- a/tests/nodes/human_input/test_human_input_node.py +++ b/tests/nodes/human_input/test_human_input_node.py @@ -6,8 +6,8 @@ import pytest +from graphon.engine_events.node import NodeRunPauseRequestedEvent, NodeRunSucceededEvent from graphon.entities.pause_reason import HitlRequired -from graphon.graph_events.node import NodeRunPauseRequestedEvent, NodeRunSucceededEvent from graphon.nodes.human_input.entities import ( Completed, Expired, @@ -17,7 +17,7 @@ PauseRequested, ) from graphon.nodes.human_input.human_input_node import HumanInputNode -from graphon.runtime.graph_runtime_state import GraphRuntimeState +from graphon.runtime.graph_runtime_state import RuntimeState from graphon.variables.segments import StringSegment from ...helpers import build_graph_init_params, build_variable_pool @@ -39,7 +39,8 @@ def _build_node( graph_config={"nodes": [], "edges": []}, run_context=run_context, ), - graph_runtime_state=GraphRuntimeState( + graph_runtime_state=RuntimeState( + workflow_id="workflow", variable_pool=build_variable_pool(), start_at=perf_counter(), ), diff --git a/tests/nodes/if_else/test_if_else_node.py b/tests/nodes/if_else/test_if_else_node.py index 40cfdb29..e91d6de0 100644 --- a/tests/nodes/if_else/test_if_else_node.py +++ b/tests/nodes/if_else/test_if_else_node.py @@ -7,11 +7,11 @@ import pytest +from graphon.engine_events.node import NodeRunSucceededEvent from graphon.enums import WorkflowNodeExecutionStatus -from graphon.graph_events.node import NodeRunSucceededEvent from graphon.node_events.base import NodeRunResult from graphon.nodes.if_else.if_else_node import IfElseNode -from graphon.runtime.graph_runtime_state import GraphRuntimeState +from graphon.runtime.graph_runtime_state import RuntimeState from ...helpers import build_graph_init_params, build_variable_pool @@ -203,7 +203,8 @@ def _run_if_else_node( data: dict[str, Any], variables: tuple[tuple[tuple[str, ...], Any], ...], ) -> tuple[NodeRunResult, list[warnings.WarningMessage]]: - runtime_state = GraphRuntimeState( + runtime_state = RuntimeState( + workflow_id="workflow", variable_pool=build_variable_pool(variables=variables), start_at=perf_counter(), ) diff --git a/tests/nodes/llm/test_node.py b/tests/nodes/llm/test_node.py index 11aca21d..9aa74fdb 100644 --- a/tests/nodes/llm/test_node.py +++ b/tests/nodes/llm/test_node.py @@ -8,15 +8,15 @@ import pytest +from graphon.engine_events.node import ( + NodeRunModelPollingProgressEvent, + NodeRunReasoningChunkEvent, +) from graphon.entities.base_node_data import BaseNodeData from graphon.enums import WorkflowNodeExecutionStatus from graphon.file import helpers as file_helpers from graphon.file.enums import FileTransferMethod, FileType from graphon.file.models import File -from graphon.graph_events.node import ( - NodeRunModelPollingProgressEvent, - NodeRunReasoningChunkEvent, -) from graphon.model_runtime.entities.llm_entities import ( LLMPollingConfig, LLMPollingResult, @@ -35,7 +35,7 @@ VideoPromptMessageContent, ) from graphon.model_runtime.entities.model_entities import ModelFeature -from graphon.node_events.base import NodeEventBase +from graphon.node_events.base import NodeEventPayload from graphon.node_events.node import ( ModelInvokeCompletedEvent, ModelPollingProgressEvent, @@ -47,7 +47,7 @@ from graphon.nodes.llm.exc import LLMNodeError from graphon.nodes.llm.reasoning import split_reasoning from graphon.nodes.llm.runtime_protocols import LLMPollingCapableProtocol, LLMProtocol -from graphon.runtime.graph_runtime_state import GraphRuntimeState +from graphon.runtime.graph_runtime_state import RuntimeState from ...helpers import build_graph_init_params, build_variable_pool @@ -166,7 +166,8 @@ def _build_llm_node( graph_config={"nodes": [], "edges": []}, run_context=run_context, ), - graph_runtime_state=GraphRuntimeState( + graph_runtime_state=RuntimeState( + workflow_id="workflow", variable_pool=build_variable_pool(variables=prepared_variables), start_at=0.0, ), @@ -495,7 +496,7 @@ def test_run_does_not_reuse_file_outputs_after_failure( def invoke( **kwargs: Any, ) -> Generator[ - NodeEventBase | LLMStructuredOutput, + NodeEventPayload | LLMStructuredOutput, None, None, ]: @@ -533,12 +534,12 @@ def invoke( ("invoke_events", "expected_error"), [ ([], "without a completion event"), - ([NodeEventBase()], "Unexpected LLM invocation event: NodeEventBase"), + ([NodeEventPayload()], "Unexpected LLM invocation event: NodeEventPayload"), ], ) def test_run_rejects_invalid_invocation_event_sequence( monkeypatch: pytest.MonkeyPatch, - invoke_events: list[NodeEventBase], + invoke_events: list[NodeEventPayload], expected_error: str, ) -> None: node = _build_llm_node() @@ -1079,7 +1080,7 @@ def _collect_stream_events( parts: Sequence[str], *, reasoning_format: Literal["separated", "tagged"], -) -> list[NodeEventBase]: +) -> list[NodeEventPayload]: """Stream ``parts`` through the LLM node and return every emitted event.""" model = MagicMock(is_structured_output_parse_error=lambda _error: False) return [ @@ -1092,7 +1093,7 @@ def _collect_stream_events( model_instance=cast(LLMProtocol, model), reasoning_format=reasoning_format, ) - if isinstance(event, NodeEventBase) + if isinstance(event, NodeEventPayload) ] diff --git a/tests/nodes/parameter_extractor/test_prompts.py b/tests/nodes/parameter_extractor/test_prompts.py index fea8b544..146c262a 100644 --- a/tests/nodes/parameter_extractor/test_prompts.py +++ b/tests/nodes/parameter_extractor/test_prompts.py @@ -25,7 +25,7 @@ FUNCTION_CALLING_EXTRACTOR_SYSTEM_PROMPT, FUNCTION_CALLING_EXTRACTOR_USER_TEMPLATE, ) -from graphon.runtime.graph_runtime_state import GraphRuntimeState +from graphon.runtime.graph_runtime_state import RuntimeState from graphon.runtime.variable_pool import VariablePool from graphon.variables.types import SegmentType @@ -44,7 +44,8 @@ def _build_dependencies(**kwargs: object) -> object: def _build_parameter_extractor_node() -> tuple[ParameterExtractorNode, VariablePool]: variable_pool = build_variable_pool(variables=[(("start", "rule"), "strictly")]) - runtime_state = GraphRuntimeState( + runtime_state = RuntimeState( + workflow_id="workflow", variable_pool=variable_pool, start_at=time.perf_counter(), ) @@ -219,7 +220,8 @@ def test_parameter_extractor_run_emits_model_identity_in_inputs( def test_parameter_extractor_accepts_dependency_bundle() -> None: variable_pool = build_variable_pool(variables=[]) - runtime_state = GraphRuntimeState( + runtime_state = RuntimeState( + workflow_id="workflow", variable_pool=variable_pool, start_at=time.perf_counter(), ) @@ -274,7 +276,8 @@ def test_parameter_extractor_rejects_mixed_dependency_styles() -> None: graph_init_params=build_graph_init_params( graph_config={"nodes": [], "edges": []} ), - graph_runtime_state=GraphRuntimeState( + graph_runtime_state=RuntimeState( + workflow_id="workflow", variable_pool=build_variable_pool(variables=[]), start_at=time.perf_counter(), ), @@ -289,7 +292,8 @@ def test_parameter_extractor_rejects_mixed_dependency_styles() -> None: def test_parameter_extractor_legacy_dependency_keywords_still_work() -> None: variable_pool = build_variable_pool(variables=[]) - runtime_state = GraphRuntimeState( + runtime_state = RuntimeState( + workflow_id="workflow", variable_pool=variable_pool, start_at=time.perf_counter(), ) diff --git a/tests/nodes/question_classifier/test_question_classifier_node.py b/tests/nodes/question_classifier/test_question_classifier_node.py index 7706dc76..2aee3538 100644 --- a/tests/nodes/question_classifier/test_question_classifier_node.py +++ b/tests/nodes/question_classifier/test_question_classifier_node.py @@ -13,7 +13,7 @@ QuestionClassifierNodeDependencies, ) from graphon.nodes.question_classifier.question_classifier_node import llm_utils -from graphon.runtime.graph_runtime_state import GraphRuntimeState +from graphon.runtime.graph_runtime_state import RuntimeState from ...helpers import build_graph_init_params @@ -78,7 +78,8 @@ def _build_question_classifier_node( graph_init_params=build_graph_init_params( graph_config={"nodes": [], "edges": []}, ), - graph_runtime_state=GraphRuntimeState( + graph_runtime_state=RuntimeState( + workflow_id="workflow", variable_pool=variable_pool, start_at=0.0, ), @@ -137,7 +138,8 @@ def test_question_classifier_constructor_accepts_dependency_bundle( graph_init_params=build_graph_init_params( graph_config={"nodes": [], "edges": []}, ), - graph_runtime_state=GraphRuntimeState( + graph_runtime_state=RuntimeState( + workflow_id="workflow", variable_pool=variable_pool, start_at=0.0, ), @@ -207,7 +209,8 @@ def test_question_classifier_constructor_rejects_mixed_dependency_inputs() -> No graph_init_params=build_graph_init_params( graph_config={"nodes": [], "edges": []}, ), - graph_runtime_state=GraphRuntimeState( + graph_runtime_state=RuntimeState( + workflow_id="workflow", variable_pool=variable_pool, start_at=0.0, ), diff --git a/tests/nodes/test_human_input_runtime_binding.py b/tests/nodes/test_human_input_runtime_binding.py index 4824ae35..fdab9bb2 100644 --- a/tests/nodes/test_human_input_runtime_binding.py +++ b/tests/nodes/test_human_input_runtime_binding.py @@ -10,7 +10,7 @@ PauseRequested, ) from graphon.nodes.human_input.human_input_node import HumanInputNode -from graphon.runtime.graph_runtime_state import GraphRuntimeState +from graphon.runtime.graph_runtime_state import RuntimeState from ..helpers import build_graph_init_params, build_variable_pool @@ -35,7 +35,8 @@ def _build_human_input_node( graph_config={"nodes": [], "edges": []}, run_context={"workflow_execution_id": "workflow-exec-1"}, ), - graph_runtime_state=GraphRuntimeState( + graph_runtime_state=RuntimeState( + workflow_id="workflow", variable_pool=build_variable_pool(), start_at=perf_counter(), ), diff --git a/tests/nodes/tool/test_tool_node.py b/tests/nodes/tool/test_tool_node.py index f4e49caf..dd6ab427 100644 --- a/tests/nodes/tool/test_tool_node.py +++ b/tests/nodes/tool/test_tool_node.py @@ -6,14 +6,14 @@ import pytest -from graphon.enums import BuiltinNodeTypes -from graphon.file.enums import FileTransferMethod, FileType -from graphon.file.models import File -from graphon.graph_events.node import ( +from graphon.engine_events.node import ( NodeRunFailedEvent, NodeRunStartedEvent, NodeRunSucceededEvent, ) +from graphon.enums import BuiltinNodeTypes +from graphon.file.enums import FileTransferMethod, FileType +from graphon.file.models import File from graphon.model_runtime.entities.llm_entities import LLMUsage from graphon.node_events.node import StreamChunkEvent, StreamCompletedEvent from graphon.nodes.tool.entities import ToolNodeData, ToolProviderType @@ -24,7 +24,7 @@ ToolRuntimeMessage, ToolRuntimeParameter, ) -from graphon.runtime.graph_runtime_state import GraphRuntimeState +from graphon.runtime.graph_runtime_state import RuntimeState from tests.helpers.builders import build_graph_init_params, build_variable_pool @@ -144,7 +144,8 @@ def _build_tool_node() -> tuple[ToolNode, _StubToolRuntime, _StubToolFileManager node_id="node-1", data=_tool_node_data(), graph_init_params=build_graph_init_params(), - graph_runtime_state=GraphRuntimeState( + graph_runtime_state=RuntimeState( + workflow_id="workflow", variable_pool=build_variable_pool(), start_at=time(), ), @@ -161,7 +162,9 @@ def _build_run_tool_node( tool_node_version: str | None, ) -> tuple[ToolNode, object]: variable_pool = build_variable_pool(variables=[(["upstream", "answer"], "42")]) - runtime_state = GraphRuntimeState(variable_pool=variable_pool, start_at=time()) + runtime_state = RuntimeState( + workflow_id="workflow", variable_pool=variable_pool, start_at=time() + ) node = ToolNode( node_id="node-1", data=_tool_node_data(tool_node_version=tool_node_version), diff --git a/tests/nodes/variable_assigner/test_v1_node.py b/tests/nodes/variable_assigner/test_v1_node.py index 24a732f0..e171b6a2 100644 --- a/tests/nodes/variable_assigner/test_v1_node.py +++ b/tests/nodes/variable_assigner/test_v1_node.py @@ -1,11 +1,14 @@ import time from collections.abc import Sequence -from graphon.graph_events.node import NodeRunSucceededEvent, NodeRunVariableUpdatedEvent +from graphon.engine_events.node import ( + NodeRunSucceededEvent, + NodeRunVariableUpdatedEvent, +) from graphon.nodes.variable_assigner.common import helpers as common_helpers from graphon.nodes.variable_assigner.v1.node import VariableAssignerNode from graphon.nodes.variable_assigner.v1.node_data import VariableAssignerData, WriteMode -from graphon.runtime.graph_runtime_state import GraphRuntimeState +from graphon.runtime.graph_runtime_state import RuntimeState from graphon.runtime.variable_pool import VariablePool from graphon.variables.variables import ( ArrayStringVariable, @@ -24,7 +27,8 @@ def _build_node( ) -> VariableAssignerNode: graph_config = {"nodes": [], "edges": []} init_params = build_graph_init_params(graph_config=graph_config) - runtime_state = GraphRuntimeState( + runtime_state = RuntimeState( + workflow_id="workflow", variable_pool=variable_pool, start_at=time.perf_counter(), ) diff --git a/tests/nodes/variable_assigner/test_v2_node.py b/tests/nodes/variable_assigner/test_v2_node.py index 6fb925e1..07ea4790 100644 --- a/tests/nodes/variable_assigner/test_v2_node.py +++ b/tests/nodes/variable_assigner/test_v2_node.py @@ -1,19 +1,19 @@ import time from collections.abc import Sequence -from graphon.enums import WorkflowNodeExecutionStatus -from graphon.graph_events.node import ( +from graphon.engine_events.node import ( NodeRunFailedEvent, NodeRunSucceededEvent, NodeRunVariableUpdatedEvent, ) +from graphon.enums import WorkflowNodeExecutionStatus from graphon.nodes.variable_assigner.v2.entities import ( VariableAssignerNodeData, VariableOperationItem, ) from graphon.nodes.variable_assigner.v2.enums import InputType, Operation from graphon.nodes.variable_assigner.v2.node import VariableAssignerNode -from graphon.runtime.graph_runtime_state import GraphRuntimeState +from graphon.runtime.graph_runtime_state import RuntimeState from graphon.runtime.variable_pool import VariablePool from graphon.variables.variables import ( ArrayStringVariable, @@ -30,7 +30,8 @@ def _build_node( ) -> VariableAssignerNode: graph_config = {"nodes": [], "edges": []} init_params = build_graph_init_params(graph_config=graph_config) - runtime_state = GraphRuntimeState( + runtime_state = RuntimeState( + workflow_id="workflow", variable_pool=variable_pool, start_at=time.perf_counter(), ) diff --git a/tests/runtime/test_custom_container_state.py b/tests/runtime/test_custom_container_state.py new file mode 100644 index 00000000..d5af22cb --- /dev/null +++ b/tests/runtime/test_custom_container_state.py @@ -0,0 +1,83 @@ +from datetime import UTC, datetime + +import pytest +from pydantic import TypeAdapter, ValidationError + +from graphon.model_runtime.entities.llm_entities import LLMUsage +from graphon.nodes.container_effects import ( + ContainerAwaitRequest, + ContainerRunResult, + CustomContainerRequest, +) +from graphon.runtime.container_state import ( + CustomContainerFrameState, + CustomContainerRunState, + FrameRuntimeData, + create_container_run_state, +) +from graphon.runtime.graph_runtime_state import RuntimeState +from graphon.runtime.variable_pool import VariablePool + + +def test_create_custom_container_run_state() -> None: + started_at = datetime.now(UTC).replace(tzinfo=None) + request = TypeAdapter(ContainerAwaitRequest).validate_python({ + "kind": "custom", + "payload": '{"version":1,"graph":"workflow-1"}', + }) + assert isinstance(request, CustomContainerRequest) + + run_state = create_container_run_state( + invocation_id="invocation-1", + frame_id="root", + node_id="workflow-tool", + started_at=started_at, + request=request, + ) + + assert run_state == CustomContainerRunState( + invocation_id="invocation-1", + frame_id="root", + node_id="workflow-tool", + started_at=started_at, + payload=request.payload, + ) + + +def test_custom_container_state_round_trips_in_runtime_snapshot() -> None: + state = RuntimeState( + workflow_id="workflow", variable_pool=VariablePool(), start_at=1 + ) + run_state = CustomContainerRunState( + invocation_id="invocation-1", + frame_id="root", + node_id="workflow-tool", + started_at=datetime.now(UTC).replace(tzinfo=None), + payload='{"version":1,"graph":"workflow-1"}', + ) + frame_state = CustomContainerFrameState( + frame_id="workflow-tool-frame", + parent_invocation_id=run_state.invocation_id, + runtime_data=FrameRuntimeData( + variable_pool=VariablePool(), + outputs={"result": "done"}, + llm_usage=LLMUsage.empty_usage(), + node_run_steps=2, + graph_node_states={}, + graph_edge_states={}, + ), + ) + state.put_container_run(run_state) + state.put_container_frame(frame_state) + + restored = RuntimeState.from_snapshot(state.dumps()) + + assert restored.get_container_run(run_state.invocation_id) == run_state + assert restored.get_container_frame(frame_state.frame_id) == frame_state + + +def test_custom_container_request_is_not_a_container_run_result() -> None: + request = CustomContainerRequest(payload="{}") + + with pytest.raises(ValidationError): + TypeAdapter(ContainerRunResult).validate_python(request.model_dump()) diff --git a/tests/runtime/test_graph_runtime_state.py b/tests/runtime/test_graph_runtime_state.py index 6c7a5f2e..f482eae0 100644 --- a/tests/runtime/test_graph_runtime_state.py +++ b/tests/runtime/test_graph_runtime_state.py @@ -5,16 +5,14 @@ import pytest -from graphon.enums import ErrorHandleMode, NodeState -from graphon.file import File, FileTransferMethod, FileType -from graphon.graph_engine.domain.graph_execution import GraphExecution -from graphon.graph_engine.ready_queue.in_memory import InMemoryReadyQueue -from graphon.graph_engine.ready_queue.protocol import ( - ROOT_FRAME_ID, +from graphon.engine.ready_queue.entities import ( ReadyTask, ResumeTask, StartTask, ) +from graphon.engine.ready_queue.in_memory import InMemoryReadyQueue +from graphon.enums import ErrorHandleMode, NodeState +from graphon.file import File, FileTransferMethod, FileType from graphon.model_runtime.entities.llm_entities import LLMUsage from graphon.nodes.container_effects import ( IterationFrameRequest, @@ -25,7 +23,8 @@ IterationFrameState, IterationRunState, ) -from graphon.runtime.graph_runtime_state import GraphRuntimeState +from graphon.runtime.execution import ROOT_FRAME_ID, GraphExecution +from graphon.runtime.graph_runtime_state import RuntimeState from graphon.runtime.read_only_wrappers import ReadOnlyGraphRuntimeStateWrapper from graphon.runtime.ready_queue import ReadyQueue from graphon.runtime.variable_pool import VariablePool @@ -64,9 +63,49 @@ def loads(self, data: str) -> None: self._queue.loads(data.removeprefix("prefixed:")) +def test_graph_execution_supplies_existing_workflow_identity() -> None: + """A child runtime preserves the aggregate shared by its parent frame.""" + execution = GraphExecution(workflow_id="workflow") + + state = RuntimeState( + variable_pool=VariablePool(), + start_at=time(), + graph_execution=execution, + ) + + assert state.graph_execution is execution + + +def test_runtime_state_requires_workflow_identity() -> None: + """Construction cannot silently create an anonymous execution aggregate.""" + with pytest.raises( + ValueError, + match="workflow_id or graph_execution is required", + ): + RuntimeState(variable_pool=VariablePool(), start_at=time()) + + +def test_runtime_state_rejects_conflicting_workflow_identity() -> None: + """Two explicit identity sources must name the same workflow.""" + execution = GraphExecution(workflow_id="existing-workflow") + + with pytest.raises( + ValueError, + match=r"workflow_id must match graph_execution\.workflow_id", + ): + RuntimeState( + workflow_id="different-workflow", + graph_execution=execution, + variable_pool=VariablePool(), + start_at=time(), + ) + + class TestGraphRuntimeState: def test_execution_context_defaults_to_empty_context(self) -> None: - state = GraphRuntimeState(variable_pool=VariablePool(), start_at=time()) + state = RuntimeState( + workflow_id="workflow", variable_pool=VariablePool(), start_at=time() + ) with state.execution_context: assert state.execution_context is not None @@ -75,7 +114,9 @@ def test_property_getters(self) -> None: variable_pool = VariablePool() start_time = time() - state = GraphRuntimeState(variable_pool=variable_pool, start_at=start_time) + state = RuntimeState( + workflow_id="workflow", variable_pool=variable_pool, start_at=start_time + ) assert state.variable_pool == variable_pool assert state.start_at == start_time @@ -83,7 +124,9 @@ def test_property_getters(self) -> None: assert state.node_run_steps == 0 def test_outputs_immutability(self) -> None: - state = GraphRuntimeState(variable_pool=VariablePool(), start_at=time()) + state = RuntimeState( + workflow_id="workflow", variable_pool=VariablePool(), start_at=time() + ) outputs1 = state.outputs outputs2 = state.outputs @@ -98,7 +141,9 @@ def test_outputs_immutability(self) -> None: assert state.get_output("key1") == "value1" def test_merge_response_outputs_appends_answer_and_overwrites_others(self) -> None: - state = GraphRuntimeState(variable_pool=VariablePool(), start_at=time()) + state = RuntimeState( + workflow_id="workflow", variable_pool=VariablePool(), start_at=time() + ) state.merge_response_outputs({"answer": "Hello", "status": "draft"}) state.merge_response_outputs({"answer": " world", "status": "final"}) @@ -107,7 +152,9 @@ def test_merge_response_outputs_appends_answer_and_overwrites_others(self) -> No assert state.get_output("status") == "final" def test_llm_usage_immutability(self) -> None: - state = GraphRuntimeState(variable_pool=VariablePool(), start_at=time()) + state = RuntimeState( + workflow_id="workflow", variable_pool=VariablePool(), start_at=time() + ) usage1 = state.llm_usage usage2 = state.llm_usage @@ -115,14 +162,17 @@ def test_llm_usage_immutability(self) -> None: def test_type_validation(self) -> None: with pytest.raises(ValueError, match="node_run_steps must be non-negative"): - GraphRuntimeState( + RuntimeState( + workflow_id="workflow", variable_pool=VariablePool(), start_at=time(), node_run_steps=-1, ) def test_helper_methods(self) -> None: - state = GraphRuntimeState(variable_pool=VariablePool(), start_at=time()) + state = RuntimeState( + workflow_id="workflow", variable_pool=VariablePool(), start_at=time() + ) initial_steps = state.node_run_steps state.increment_node_run_steps() @@ -134,20 +184,24 @@ def test_helper_methods(self) -> None: assert state.llm_usage.total_tokens == 50 def test_ready_queue_default_instantiation(self) -> None: - state = GraphRuntimeState(variable_pool=VariablePool(), start_at=time()) + state = RuntimeState( + workflow_id="workflow", variable_pool=VariablePool(), start_at=time() + ) queue = state.ready_queue assert isinstance(queue, InMemoryReadyQueue) def test_deferred_ready_tasks_round_trip_in_runtime_snapshot(self) -> None: - state = GraphRuntimeState(variable_pool=VariablePool(), start_at=time()) + state = RuntimeState( + workflow_id="workflow", variable_pool=VariablePool(), start_at=time() + ) first = StartTask(frame_id="root", node_id="a") second = StartTask(frame_id="child", node_id="b") state.defer_ready_task(first) state.defer_ready_task(second) - restored = GraphRuntimeState.from_snapshot(state.dumps()) + restored = RuntimeState.from_snapshot(state.dumps()) assert restored.drain_deferred_ready_tasks() == [first, second] assert restored.drain_deferred_ready_tasks() == [] @@ -155,7 +209,8 @@ def test_deferred_ready_tasks_round_trip_in_runtime_snapshot(self) -> None: def test_custom_ready_queues_round_trip_with_supplied_factory(self) -> None: ready_queue: ReadyQueue = _PrefixedReadyQueue() deferred_ready_queue: ReadyQueue = _PrefixedReadyQueue() - state = GraphRuntimeState( + state = RuntimeState( + workflow_id="workflow", variable_pool=VariablePool(), start_at=time(), ready_queue=ready_queue, @@ -166,7 +221,7 @@ def test_custom_ready_queues_round_trip_with_supplied_factory(self) -> None: state.ready_queue.put(live_task) state.defer_ready_task(deferred_task) - restored = GraphRuntimeState.from_snapshot( + restored = RuntimeState.from_snapshot( state.dumps(), ready_queue_factory=_PrefixedReadyQueue, ) @@ -176,7 +231,9 @@ def test_custom_ready_queues_round_trip_with_supplied_factory(self) -> None: assert restored.drain_deferred_ready_tasks() == [deferred_task] def test_container_runtime_state_preserves_file_values(self) -> None: - state = GraphRuntimeState(variable_pool=VariablePool(), start_at=time()) + state = RuntimeState( + workflow_id="workflow", variable_pool=VariablePool(), start_at=time() + ) file_value = File( file_id="file-1", file_type=FileType.DOCUMENT, @@ -229,7 +286,7 @@ def test_container_runtime_state_preserves_file_values(self) -> None: ResumeTask(invocation_id=run.invocation_id, result=request), ) - restored = GraphRuntimeState.from_snapshot(state.dumps()) + restored = RuntimeState.from_snapshot(state.dumps()) restored_run = restored.get_container_run("invocation-1") assert isinstance(restored_run, IterationRunState) @@ -242,17 +299,24 @@ def test_container_runtime_state_preserves_file_values(self) -> None: assert restored_task.result.items[0].value == file_value assert restored.get_container_frame("exec-iteration:iteration:0") == frame - def test_graph_execution_lazy_instantiation(self) -> None: - state = GraphRuntimeState(variable_pool=VariablePool(), start_at=time()) + def test_workflow_id_creates_graph_execution(self) -> None: + """A root runtime creates its execution aggregate from its workflow ID.""" + state = RuntimeState( + workflow_id="workflow", + variable_pool=VariablePool(), + start_at=time(), + ) execution = state.graph_execution assert isinstance(execution, GraphExecution) - assert not execution.workflow_id + assert execution.workflow_id == "workflow" assert state.graph_execution is execution def test_graph_configuration_rejects_different_graph(self) -> None: - state = GraphRuntimeState(variable_pool=VariablePool(), start_at=time()) + state = RuntimeState( + workflow_id="workflow", variable_pool=VariablePool(), start_at=time() + ) mock_graph = MagicMock() state.attach_graph(mock_graph) @@ -261,19 +325,41 @@ def test_graph_configuration_rejects_different_graph(self) -> None: other_graph = MagicMock() with pytest.raises( ValueError, - match="GraphRuntimeState already attached to a different graph instance", + match="RuntimeState already attached to a different graph instance", ): state.attach_graph(other_graph) + def test_attach_graph_rejects_a_different_materialized_scope(self) -> None: + state = RuntimeState( + workflow_id="workflow", variable_pool=VariablePool(), start_at=time() + ) + state.restore_graph_state( + node_states={ + "root": NodeState.TAKEN, + "child": NodeState.SKIPPED, + }, + edge_states={"child-edge": NodeState.SKIPPED}, + ) + graph = MagicMock(nodes={"root": MagicMock()}, edges={}) + + with pytest.raises( + RuntimeError, + match="Saved graph state does not match rebuilt graph", + ): + state.attach_graph(graph) + def test_read_only_wrapper_exposes_additional_state(self) -> None: - state = GraphRuntimeState(variable_pool=VariablePool(), start_at=time()) + state = RuntimeState( + workflow_id="workflow", variable_pool=VariablePool(), start_at=time() + ) wrapper = ReadOnlyGraphRuntimeStateWrapper(state) assert wrapper.ready_queue_size == 0 assert wrapper.exceptions_count == 0 def test_read_only_wrapper_serializes_runtime_state(self) -> None: - state = GraphRuntimeState( + state = RuntimeState( + workflow_id="workflow", variable_pool=VariablePool(), start_at=time(), llm_usage=LLMUsage.from_metadata({"total_tokens": 5}), @@ -300,7 +386,8 @@ def test_dumps_and_loads_roundtrip(self) -> None: "currency": "USD", "latency": 0.5, }) - state = GraphRuntimeState( + state = RuntimeState( + workflow_id="wf-123", variable_pool=variable_pool, start_at=time(), node_run_steps=3, @@ -310,14 +397,13 @@ def test_dumps_and_loads_roundtrip(self) -> None: state.ready_queue.put(StartTask(frame_id="root", node_id="node-A")) graph_execution = state.graph_execution - graph_execution.workflow_id = "wf-123" graph_execution.exceptions_count = 4 graph_execution.started = True graph_execution.error = ValueError("saved failure") snapshot = state.dumps() - restored = GraphRuntimeState.from_snapshot(snapshot) + restored = RuntimeState.from_snapshot(snapshot) assert restored.total_tokens == 5 assert restored.node_run_steps == 3 @@ -386,10 +472,10 @@ def test_version_1_snapshot_migrates_to_frame_aware_version_2(self) -> None: }, }) - restored = GraphRuntimeState.from_snapshot(snapshot) + restored = RuntimeState.from_snapshot(snapshot) migrated = json.loads(restored.dumps()) - assert migrated["version"] == "2.0" + assert migrated["version"] == "3.0" assert json.loads(migrated["ready_queue"])["version"] == "2.0" assert json.loads(migrated["deferred_ready_tasks"])["version"] == "2.0" assert json.loads(migrated["graph_execution"])["version"] == "2.0" @@ -423,9 +509,11 @@ def test_snapshot_restore_preserves_updated_conversation_variable(self) -> None: ) variable_pool.add((CONVERSATION_VARIABLE_NODE_ID, "session_name"), "after") - state = GraphRuntimeState(variable_pool=variable_pool, start_at=time()) + state = RuntimeState( + workflow_id="workflow", variable_pool=variable_pool, start_at=time() + ) snapshot = state.dumps() - restored = GraphRuntimeState.from_snapshot(snapshot) + restored = RuntimeState.from_snapshot(snapshot) restored_value = restored.variable_pool.get(( CONVERSATION_VARIABLE_NODE_ID, @@ -449,9 +537,11 @@ def test_snapshot_restore_preserves_file_segments(self) -> None: variable_pool.add(("node", "attachment"), FileSegment(value=file_value)) variable_pool.add(("node", "attachments"), ArrayFileSegment(value=[file_value])) - state = GraphRuntimeState(variable_pool=variable_pool, start_at=time()) + state = RuntimeState( + workflow_id="workflow", variable_pool=variable_pool, start_at=time() + ) - restored = GraphRuntimeState.from_snapshot(state.dumps()) + restored = RuntimeState.from_snapshot(state.dumps()) restored_file = restored.variable_pool.get(("node", "attachment")) restored_files = restored.variable_pool.get(("node", "attachments")) diff --git a/tests/test_protocol_abstract_contracts.py b/tests/test_protocol_abstract_contracts.py deleted file mode 100644 index 8537f8be..00000000 --- a/tests/test_protocol_abstract_contracts.py +++ /dev/null @@ -1,132 +0,0 @@ -from __future__ import annotations - -import ast -import importlib -import inspect -from pathlib import Path - -import pytest - - -def _is_direct_protocol_class( - class_def: ast.ClassDef, - *, - protocol_aliases: set[str], -) -> bool: - for base in class_def.bases: - base_expr = base.value if isinstance(base, ast.Subscript) else base - if isinstance(base_expr, ast.Name) and base_expr.id in protocol_aliases: - return True - if ( - isinstance(base_expr, ast.Attribute) - and isinstance(base_expr.value, ast.Name) - and f"{base_expr.value.id}.{base_expr.attr}" in protocol_aliases - ): - return True - return False - - -def _discover_protocol_aliases(parsed: ast.Module) -> set[str]: - protocol_aliases = set[str]() - typing_aliases: set[str] = set() - - for node in parsed.body: - if isinstance(node, ast.ImportFrom) and node.module in { - "typing", - "typing_extensions", - }: - for alias in node.names: - if alias.name == "Protocol": - protocol_aliases.add(alias.asname or alias.name) - continue - - if isinstance(node, ast.Import): - for alias in node.names: - if alias.name in {"typing", "typing_extensions"}: - typing_aliases.add(alias.asname or alias.name) - - protocol_aliases.update(f"{alias}.Protocol" for alias in typing_aliases) - return protocol_aliases - - -def _has_protocol_members(class_def: ast.ClassDef) -> bool: - return any( - isinstance(member, ast.FunctionDef | ast.AsyncFunctionDef) - for member in class_def.body - ) - - -def _discover_protocol_targets() -> list[type[object]]: - src_root = Path(__file__).resolve().parents[1] / "src" / "graphon" - protocol_classes: list[type[object]] = [] - - for file_path in sorted(src_root.rglob("*.py")): - parsed = ast.parse(file_path.read_text()) - protocol_aliases = _discover_protocol_aliases(parsed) - if not protocol_aliases: - continue - - module_name = "graphon." + ".".join( - file_path.relative_to(src_root).with_suffix("").parts, - ) - module = importlib.import_module(module_name) - - for class_def in [ - node for node in parsed.body if isinstance(node, ast.ClassDef) - ]: - if not _is_direct_protocol_class( - class_def, - protocol_aliases=protocol_aliases, - ): - continue - if not _has_protocol_members(class_def): - continue - protocol_classes.append(getattr(module, class_def.name)) - - protocol_classes.sort(key=lambda cls: (cls.__module__, cls.__qualname__)) - return protocol_classes - - -def _protocol_member_names(protocol_cls: type[object]) -> list[str]: - member_names: list[str] = [] - for name, value in protocol_cls.__dict__.items(): - if name.startswith("__") and name.endswith("__"): - continue - if isinstance(value, property | classmethod | staticmethod): - member_names.append(name) - continue - if inspect.isfunction(value): - member_names.append(name) - return member_names - - -PROTOCOL_TARGETS = _discover_protocol_targets() - - -def _protocol_id(protocol_cls: type[object]) -> str: - return f"{protocol_cls.__module__}.{protocol_cls.__qualname__}" - - -def test_protocol_targets_should_be_discovered() -> None: - assert PROTOCOL_TARGETS - - -@pytest.mark.parametrize( - "protocol_cls", - PROTOCOL_TARGETS, - ids=_protocol_id, -) -def test_protocol_members_should_be_abstract(protocol_cls: type[object]) -> None: - member_names = _protocol_member_names(protocol_cls) - # This protects the test from a vacuous pass if discovery and runtime member - # detection drift apart. Discovery only targets Protocol classes with - # methods, so an empty member list means the test is no longer checking the - # contract it claims to check. - assert member_names, f"{_protocol_id(protocol_cls)} has no protocol members." - - non_abstract_members = [ - name - for name in member_names - if not getattr(protocol_cls.__dict__[name], "__isabstractmethod__", False) - ] - assert non_abstract_members == [] diff --git a/tests/workflows/test_full_graph_events.py b/tests/workflows/test_full_engine_events.py similarity index 89% rename from tests/workflows/test_full_graph_events.py rename to tests/workflows/test_full_engine_events.py index d251b87a..9cc9ce42 100644 --- a/tests/workflows/test_full_graph_events.py +++ b/tests/workflows/test_full_engine_events.py @@ -8,27 +8,26 @@ import yaml from graphon.dsl import loads -from graphon.graph_engine.config import GraphEngineConfig -from graphon.graph_engine.frames import FrameRegistry -from graphon.graph_engine.graph_engine import GraphEngine -from graphon.graph_engine.loop_container_handler import LoopContainerHandler -from graphon.graph_engine.ready_queue.in_memory import InMemoryReadyQueue -from graphon.graph_engine.ready_queue.protocol import StartTask -from graphon.graph_events.base import GraphEngineEvent -from graphon.graph_events.graph import ( +from graphon.engine import Engine +from graphon.engine.container_handler import LoopContainerHandler +from graphon.engine.frame import FrameRegistry +from graphon.engine.ready_queue.entities import StartTask +from graphon.engine.ready_queue.in_memory import InMemoryReadyQueue +from graphon.engine_events.base import EngineEvent +from graphon.engine_events.graph import ( GraphRunFailedEvent, GraphRunPartialSucceededEvent, ) -from graphon.graph_events.iteration import ( +from graphon.engine_events.iteration import ( NodeRunIterationFailedEvent, NodeRunIterationNextEvent, NodeRunIterationSucceededEvent, ) -from graphon.graph_events.loop import ( +from graphon.engine_events.loop import ( NodeRunLoopFailedEvent, NodeRunLoopSucceededEvent, ) -from graphon.graph_events.node import NodeRunSucceededEvent +from graphon.engine_events.node import NodeRunSucceededEvent from graphon.variables.segments import Segment from tests.helpers.workflow_events import ( event_path, @@ -123,33 +122,32 @@ def _iteration_dsl(*, is_parallel: bool, parallel_nums: int) -> str: def _event( event_type: str, subject: str = "", - in_loop: str = "", - in_iteration: str = "", -) -> tuple[str, str, str, str]: - return event_type, subject, in_loop, in_iteration + container_id: str = "", +) -> tuple[str, str, str]: + return event_type, subject, container_id def _run_failed_workflow( dsl: str, *, start_inputs: Mapping[str, object], -) -> list[GraphEngineEvent]: - events: list[GraphEngineEvent] = [] +) -> list[EngineEvent]: + events: list[EngineEvent] = [] engine = loads(dsl, start_inputs=start_inputs) with pytest.raises(RuntimeError, match="Variable"): events.extend(engine.run()) return events -def _use_bounded_ready_queue(engine: GraphEngine) -> InMemoryReadyQueue: +def _use_bounded_ready_queue(engine: Engine) -> InMemoryReadyQueue: ready_queue = InMemoryReadyQueue(maxsize=1) engine.graph_runtime_state._ready_queue = ready_queue engine._worker_pool._ready_queue = ready_queue return ready_queue -def _run_with_timeout(engine: GraphEngine) -> list[GraphEngineEvent]: - events: list[GraphEngineEvent] = [] +def _run_with_timeout(engine: Engine) -> list[EngineEvent]: + events: list[EngineEvent] = [] errors: list[Exception] = [] finished = Event() @@ -252,7 +250,7 @@ def test_resume_replays_tasks_through_a_bounded_ready_queue() -> None: ) engine = loads( dsl, - config=GraphEngineConfig(min_workers=1, max_workers=1), + workers=1, ) ready_queue = _use_bounded_ready_queue(engine) engine.graph_runtime_state.graph_execution.start() @@ -299,13 +297,21 @@ def test_full_iteration_graph_records_process_and_final_outputs() -> None: _event("NodeRunStartedEvent", "iteration"), _event("NodeRunIterationStartedEvent", "iteration"), _event("NodeRunIterationNextEvent", "iteration"), - _event("GraphEdgeTakenEvent", "iteration-start->render"), - _event("NodeRunStartedEvent", "render", in_iteration="iteration"), - _event("NodeRunSucceededEvent", "render", in_iteration="iteration"), + _event( + "GraphEdgeTakenEvent", + "iteration-start->render", + container_id="iteration", + ), + _event("NodeRunStartedEvent", "render", container_id="iteration"), + _event("NodeRunSucceededEvent", "render", container_id="iteration"), _event("NodeRunIterationNextEvent", "iteration"), - _event("GraphEdgeTakenEvent", "iteration-start->render"), - _event("NodeRunStartedEvent", "render", in_iteration="iteration"), - _event("NodeRunSucceededEvent", "render", in_iteration="iteration"), + _event( + "GraphEdgeTakenEvent", + "iteration-start->render", + container_id="iteration", + ), + _event("NodeRunStartedEvent", "render", container_id="iteration"), + _event("NodeRunSucceededEvent", "render", container_id="iteration"), _event("NodeRunIterationSucceededEvent", "iteration"), _event("GraphEdgeTakenEvent", "iteration->end"), _event("NodeRunSucceededEvent", "iteration"), @@ -330,7 +336,7 @@ def test_parallel_iteration_with_one_worker_and_a_bounded_ready_queue() -> None: engine = loads( dsl, start_inputs={"items": ["alpha", "beta"]}, - config=GraphEngineConfig(min_workers=1, max_workers=1), + workers=1, ) _use_bounded_ready_queue(engine) @@ -381,7 +387,11 @@ def test_full_loop_graph_breaks_at_the_configured_condition( }, { "id": "loop-start", - "data": {"type": "loop-start", "loop_id": "loop"}, + "data": { + "type": "loop-start", + "loop_id": "loop", + "container_id": "loop", + }, }, { "id": "increment", @@ -425,10 +435,14 @@ def test_full_loop_graph_breaks_at_the_configured_condition( ] for index in range(completed_rounds): expected_path.extend([ - _event("GraphEdgeTakenEvent", "loop-start->increment"), - _event("NodeRunStartedEvent", "increment", in_loop="loop"), - _event("NodeRunVariableUpdatedEvent", "increment", in_loop="loop"), - _event("NodeRunSucceededEvent", "increment", in_loop="loop"), + _event( + "GraphEdgeTakenEvent", + "loop-start->increment", + container_id="loop", + ), + _event("NodeRunStartedEvent", "increment", container_id="loop"), + _event("NodeRunVariableUpdatedEvent", "increment", container_id="loop"), + _event("NodeRunSucceededEvent", "increment", container_id="loop"), ]) if index + 1 < completed_rounds: expected_path.append(_event("NodeRunLoopNextEvent", "loop")) @@ -493,9 +507,13 @@ def test_full_loop_graph_stops_at_loop_end_node() -> None: _event("NodeRunSucceededEvent", "start"), _event("NodeRunStartedEvent", "loop"), _event("NodeRunLoopStartedEvent", "loop"), - _event("GraphEdgeTakenEvent", "loop-start->stop"), - _event("NodeRunStartedEvent", "stop", in_loop="loop"), - _event("NodeRunSucceededEvent", "stop", in_loop="loop"), + _event( + "GraphEdgeTakenEvent", + "loop-start->stop", + container_id="loop", + ), + _event("NodeRunStartedEvent", "stop", container_id="loop"), + _event("NodeRunSucceededEvent", "stop", container_id="loop"), _event("NodeRunLoopSucceededEvent", "loop"), _event("GraphEdgeTakenEvent", "loop->end"), _event("NodeRunSucceededEvent", "loop"), @@ -578,6 +596,7 @@ def test_nested_iteration_loop_end_stops_ancestor_loop() -> None: "data": { "type": "iteration", "loop_id": "loop", + "container_id": "loop", "iterator_selector": ["start", "items"], "output_selector": ["iteration", "item"], "start_node_id": "iteration-start", @@ -593,6 +612,7 @@ def test_nested_iteration_loop_end_stops_ancestor_loop() -> None: "type": "iteration-start", "loop_id": "loop", "iteration_id": "iteration", + "container_id": "iteration", }, }, { @@ -601,6 +621,7 @@ def test_nested_iteration_loop_end_stops_ancestor_loop() -> None: "type": "loop-end", "loop_id": "loop", "iteration_id": "iteration", + "container_id": "iteration", }, }, _end_node([]), @@ -626,8 +647,7 @@ def test_nested_iteration_loop_end_stops_ancestor_loop() -> None: iteration_succeeded = [ event for event in events if isinstance(event, NodeRunIterationSucceededEvent) ] - assert stop_succeeded.in_loop_id == "loop" - assert stop_succeeded.in_iteration_id == "iteration" + assert stop_succeeded.container_id == "iteration" assert loop_succeeded.outputs == {"loop_round": 1} assert loop_succeeded.metadata["completed_reason"] == "loop_break" assert len(iteration_succeeded) == 1 @@ -715,7 +735,7 @@ def test_full_loop_graph_propagates_child_failure() -> None: ) engine = loads(dsl, start_inputs={"value": "partial"}) - events: list[GraphEngineEvent] = [] + events: list[EngineEvent] = [] with pytest.raises(RuntimeError, match="Variable"): events.extend(engine.run()) @@ -726,12 +746,20 @@ def test_full_loop_graph_propagates_child_failure() -> None: _event("NodeRunSucceededEvent", "start"), _event("NodeRunStartedEvent", "loop"), _event("NodeRunLoopStartedEvent", "loop"), - _event("GraphEdgeTakenEvent", "loop-start->child-response"), - _event("NodeRunStartedEvent", "child-response", in_loop="loop"), - _event("GraphEdgeTakenEvent", "child-response->fail"), - _event("NodeRunSucceededEvent", "child-response", in_loop="loop"), - _event("NodeRunStartedEvent", "fail", in_loop="loop"), - _event("NodeRunFailedEvent", "fail", in_loop="loop"), + _event( + "GraphEdgeTakenEvent", + "loop-start->child-response", + container_id="loop", + ), + _event("NodeRunStartedEvent", "child-response", container_id="loop"), + _event( + "GraphEdgeTakenEvent", + "child-response->fail", + container_id="loop", + ), + _event("NodeRunSucceededEvent", "child-response", container_id="loop"), + _event("NodeRunStartedEvent", "fail", container_id="loop"), + _event("NodeRunFailedEvent", "fail", container_id="loop"), _event("NodeRunLoopFailedEvent", "loop"), _event("NodeRunFailedEvent", "loop"), _event("GraphRunFailedEvent"), @@ -915,20 +943,28 @@ def test_full_iteration_graph_applies_error_handling_mode( for _ in range(executed_items): expected_path.append(_event("NodeRunIterationNextEvent", "iteration")) expected_path.extend([ - _event("GraphEdgeTakenEvent", "iteration-start->child-response"), + _event( + "GraphEdgeTakenEvent", + "iteration-start->child-response", + container_id="iteration", + ), _event( "NodeRunStartedEvent", "child-response", - in_iteration="iteration", + container_id="iteration", + ), + _event( + "GraphEdgeTakenEvent", + "child-response->fail", + container_id="iteration", ), - _event("GraphEdgeTakenEvent", "child-response->fail"), _event( "NodeRunSucceededEvent", "child-response", - in_iteration="iteration", + container_id="iteration", ), - _event("NodeRunStartedEvent", "fail", in_iteration="iteration"), - _event("NodeRunFailedEvent", "fail", in_iteration="iteration"), + _event("NodeRunStartedEvent", "fail", container_id="iteration"), + _event("NodeRunFailedEvent", "fail", container_id="iteration"), ]) if error_handle_mode == "terminated": expected_path.extend([