fix(langgraph): wire node stream on_next handlers to registered executors
CI / push-validation (pull_request) Successful in 38s
CI / helm (pull_request) Successful in 39s
CI / lint (pull_request) Successful in 3m55s
CI / build (pull_request) Successful in 3m50s
CI / typecheck (pull_request) Successful in 4m33s
CI / quality (pull_request) Successful in 4m58s
CI / security (pull_request) Successful in 5m18s
CI / e2e_tests (pull_request) Successful in 8m59s
CI / integration_tests (pull_request) Successful in 10m48s
CI / unit_tests (pull_request) Successful in 11m51s
CI / docker (pull_request) Successful in 1m55s
CI / coverage (pull_request) Successful in 16m37s
CI / status-check (pull_request) Successful in 3s

Replace the no-op on_next handlers in _setup_node_stream_subscriptions
with factory-created callbacks that look up and invoke the per-node
sync_executor registered by _register_node_executor.  The previous code
silently discarded every message delivered to node streams.

Key changes:

- Wire on_next to executor via _make_on_next_handler factory method
- Store executors in _node_executors dict (not setattr on stream_router)
- Shared ThreadPoolExecutor with proper lifecycle (start/stop/restart)
- Prefer run_coroutine_threadsafe when scheduler loop is running;
  fall back to thread pool with deadlock prevention
- Thread-safe StateManager: Lock on all mutation methods, deep-copy
  emission inside lock for immutable subscriber snapshots
- replace_state() and reset() deep-copy input to prevent external
  mutation bypassing the lock (review cycle 5 M1 fix)
- Best-effort cancellation documented on timed-out futures (cycle 5 M2)
- CancelledError catch in thread pool path with "graph stopping" message
- is_running guard on both sync_executor and execute() entry point
- Checkpoint filenames include update_count to prevent collisions
- _save_checkpoint docstring warns against calling while _lock is held
- on_next handler distinguishes TimeoutError (warning) from other
  exceptions (exception log) for clearer diagnostics
- TODO comment for deferred downstream propagation to successor nodes
- __del__ safety net for executor pool cleanup
- File I/O outside lock for checkpoint save/load
- Bounded execution_history via collections.deque(maxlen=1000)

See PR !10795 for the full change log across 5 review cycles.

ISSUES CLOSED: #6511
This commit is contained in:
2026-04-20 07:54:57 +00:00
parent e17a6ddec7
commit 0d267934a7
5 changed files with 1224 additions and 90 deletions
+124
View File
@@ -118,6 +118,48 @@ Feature: Consolidated Langgraph
Then the subscription handlers should be exercised
Scenario: Node stream messages are processed by the node executor
Given a langgraph instance with a mocked node executor
When a message is emitted on the worker node stream
Then the registered executor should be called with the message
Scenario: Raw stream values are wrapped before node processing
Given a langgraph instance with a mocked node executor
When a raw value is emitted on the worker node stream
Then the executor should receive a StreamMessage wrapping the raw value
Scenario: Missing executor triggers a diagnostic warning
Given a langgraph instance with a removed node executor
When a message is emitted on the executorless worker stream
Then a warning about missing executor should be logged
Scenario: Each node stream dispatches to its own executor
Given a langgraph instance with multiple mocked node executors
When messages are emitted on both alpha and beta node streams
Then each executor should receive its own message independently
Scenario: Node stream on_next handles executor exceptions gracefully
Given a langgraph instance with a failing node executor
When a message is emitted on the failing worker stream
Then the exception should be caught and logged without propagating
Scenario: Subscription is skipped when observable is missing
Given a langgraph instance with a removed observable
When node stream subscriptions are set up
Then no executor should be called and no error should occur
Scenario: Execute raises when start stream is missing
Given a langgraph instance with a removed start stream for execute
When I execute the graph expecting a start stream error
Then a start stream missing error should be raised from execute
Scenario: Parallel groups computed for converging branches
Given a langgraph config with converging branches
When I construct the LangGraph for parallel groups
@@ -136,6 +178,88 @@ Feature: Consolidated Langgraph
Then it should use default state and update running flags
Scenario: Executor completes via run_coroutine_threadsafe on running loop
Given a langgraph instance with a scheduler whose loop is running in a background thread
When a message is emitted on the worker stream via the running loop path
Then the executor should complete successfully via run_coroutine_threadsafe
Scenario: Execution history is capped at maximum size
Given a langgraph instance with execution history at capacity
When one more executor invocation is recorded
Then the history length should remain at the maximum and the oldest entry should be dropped
Scenario: Concurrent state updates are serialized by the lock
Given a state manager configured for concurrent access
When multiple threads perform state updates simultaneously
Then the final execution count should equal the total number of updates
Scenario: Executor times out via run_coroutine_threadsafe path
Given a langgraph instance with a slow executor on a background loop and a short timeout
When the slow executor is invoked via the running loop path
Then a TimeoutError should be raised with the node name and timeout duration
Scenario: Executor times out via thread pool path
Given a langgraph instance with a slow executor and a short timeout
When the slow executor is invoked via the thread pool path
Then a TimeoutError should be raised from the thread pool path with the node name
Scenario: Graph can be restarted after stop
Given a langgraph instance that has been started and stopped
When I restart the graph and invoke an executor
Then the executor should complete successfully after restart
Scenario: Executor raises when graph is not running
Given a langgraph instance that is not running
When I invoke the node executor directly
Then a RuntimeError should be raised indicating the graph is not running
Scenario: StateManager replace_state atomically replaces and emits deep copy
Given a state manager with an initial state
When I call replace_state with a new state
Then get_state should return the new state
And the state stream should have received a notification
And the emitted state should be a deep copy not the same object
And the emitted state should be distinct from the internal state
Scenario: CancelledError in thread pool path raises descriptive RuntimeError
Given a langgraph instance with a cancellable executor task
When the executor future is cancelled before completion
Then a RuntimeError should be raised indicating graph stopping
Scenario: LangGraph __del__ shuts down executor pool without error
Given a langgraph instance that was never stopped
When __del__ is called on the graph
Then no exception should propagate from __del__
Scenario: Execute raises when graph is not running via execute method
Given a langgraph instance that is not running for execute
When I call execute on the non-running graph
Then a RuntimeError should be raised indicating the graph is not running for execute
Scenario: TimeoutError through stream path logs warning not exception
Given a langgraph instance with an executor that raises TimeoutError
When a message is emitted on the timeout worker stream
Then the timeout should be logged at warning level
And logger exception should not be called for timeout
Scenario: CancelledError in run_coroutine_threadsafe path raises descriptive RuntimeError
Given a langgraph instance with a scheduler loop and a cancellable coroutine future
When the coroutine future is cancelled before completion
Then a RuntimeError should be raised indicating graph stopping from coroutine path
# ============================================================
# Originally from: langgraph_nodes_additional_coverage.feature
# Feature: Additional LangGraph nodes coverage
+782 -10
View File
@@ -1,12 +1,18 @@
import asyncio
from unittest.mock import AsyncMock, MagicMock
import concurrent.futures
import threading
from unittest.mock import AsyncMock, MagicMock, patch
from behave import given, then, when
from rx.scheduler.eventloop import AsyncIOScheduler
from cleveragents.langgraph.graph import GraphConfig, LangGraph
from cleveragents.langgraph.graph import (
MAX_EXECUTION_HISTORY,
GraphConfig,
LangGraph,
)
from cleveragents.langgraph.nodes import Edge, NodeConfig, NodeType
from cleveragents.langgraph.state import GraphState
from cleveragents.langgraph.state import GraphState, StateManager
from cleveragents.reactive.stream_router import StreamMessage
# Existing steps cover constructor scheduling, execution dispatch, start validation,
@@ -42,7 +48,7 @@ def step_assert_levels(context):
def step_history_instance(context):
config = GraphConfig(name="graph-history")
context.graph = _fresh_graph(config)
context.graph.execution_history = ["x", "y"]
context.graph.execution_history.extend(["x", "y"])
@when("I request execution history")
@@ -216,6 +222,7 @@ def step_assert_scheduler_and_streams(context):
def step_graph_with_state_manager(context):
config = GraphConfig(name="graph-exec")
context.graph = _fresh_graph(config)
context.graph.is_running = True
context.start_stream = f"__{context.graph.name}_node_start__"
context.graph.stream_router.send_message = MagicMock()
@@ -224,10 +231,14 @@ def step_graph_with_state_manager(context):
def step_execute_graph(context):
loop = asyncio.new_event_loop()
asyncio.set_event_loop(loop)
result = loop.run_until_complete(context.graph.execute({"messages": []}))
loop.close()
asyncio.set_event_loop(asyncio.new_event_loop())
context.execute_result = result
try:
result = loop.run_until_complete(context.graph.execute({"messages": []}))
context.execute_result = result
finally:
loop.close()
replacement_loop = asyncio.new_event_loop()
asyncio.set_event_loop(replacement_loop)
context.add_cleanup(replacement_loop.close)
@then("the start stream should receive the state and execution returns a GraphState")
@@ -292,10 +303,11 @@ def step_graph_with_executor(context):
}
config = GraphConfig(name="graph-executor", nodes=nodes)
graph = _fresh_graph(config)
graph.is_running = True
graph.state_manager.get_state = MagicMock(return_value=GraphState())
graph.state_manager.update_state = MagicMock()
graph.nodes["worker"].execute = AsyncMock(return_value={"messages": ["m"]})
context.executor = graph.stream_router._builtin_execute_node_worker
context.executor = graph._node_executors["worker"] # pylint: disable=protected-access
context.graph = graph
@@ -308,5 +320,765 @@ def step_invoke_executor(context):
@then("it should update state and record execution history")
def step_assert_executor(context):
context.graph.state_manager.update_state.assert_called_once()
assert context.graph.execution_history == ["worker"]
# deque does not compare equal to list; convert first
assert list(context.graph.execution_history) == ["worker"]
assert isinstance(context.executor_result, StreamMessage)
# ── Node stream on_next wiring tests ──────────────────────────────────
@given("a langgraph instance with a mocked node executor")
def step_prepare_mocked_executor(context):
nodes = {"worker": NodeConfig(name="worker", type=NodeType.FUNCTION)}
context.graph = _fresh_graph(GraphConfig(name="graph-exec-wired", nodes=nodes))
context.mock_executor = MagicMock(
return_value=StreamMessage(content={}, metadata={})
)
context.graph._node_executors["worker"] = context.mock_executor # pylint: disable=protected-access
@when("a message is emitted on the worker node stream")
def step_emit_worker_stream_message(context):
worker_stream = context.graph.stream_router.streams[
f"__{context.graph.name}_node_worker__"
]
worker_stream.on_next(StreamMessage(content="test-payload", metadata={}))
@then("the registered executor should be called with the message")
def step_assert_executor_called(context):
context.mock_executor.assert_called_once()
call_args = context.mock_executor.call_args
msg = call_args[0][0]
assert isinstance(msg, StreamMessage)
assert msg.content == "test-payload"
@when("a raw value is emitted on the worker node stream")
def step_emit_raw_value_on_stream(context):
worker_stream = context.graph.stream_router.streams[
f"__{context.graph.name}_node_worker__"
]
worker_stream.on_next("raw-string-payload")
@then("the executor should receive a StreamMessage wrapping the raw value")
def step_assert_wrapped_message(context):
context.mock_executor.assert_called_once()
call_args = context.mock_executor.call_args
msg = call_args[0][0]
assert isinstance(msg, StreamMessage)
assert msg.content == "raw-string-payload"
assert msg.metadata == {}
@given("a langgraph instance with a removed node executor")
def step_prepare_missing_executor(context):
nodes = {"worker": NodeConfig(name="worker", type=NodeType.FUNCTION)}
context.graph = _fresh_graph(GraphConfig(name="graph-no-exec", nodes=nodes))
context.graph._node_executors.pop("worker", None) # pylint: disable=protected-access
context.graph.logger = MagicMock()
@when("a message is emitted on the executorless worker stream")
def step_emit_executorless_message(context):
worker_stream = context.graph.stream_router.streams[
f"__{context.graph.name}_node_worker__"
]
worker_stream.on_next("test")
@then("a warning about missing executor should be logged")
def step_assert_warning_logged(context):
context.graph.logger.warning.assert_called_once()
args, _kwargs = context.graph.logger.warning.call_args
assert args[0] == "No executor found for node %s"
assert args[1] == "worker"
@given("a langgraph instance with multiple mocked node executors")
def step_prepare_multi_executors(context):
nodes = {
"alpha": NodeConfig(name="alpha", type=NodeType.FUNCTION),
"beta": NodeConfig(name="beta", type=NodeType.FUNCTION),
}
context.graph = _fresh_graph(GraphConfig(name="graph-multi", nodes=nodes))
context.alpha_executor = MagicMock(
return_value=StreamMessage(content={}, metadata={})
)
context.beta_executor = MagicMock(
return_value=StreamMessage(content={}, metadata={})
)
context.graph._node_executors["alpha"] = context.alpha_executor # pylint: disable=protected-access
context.graph._node_executors["beta"] = context.beta_executor # pylint: disable=protected-access
@when("messages are emitted on both alpha and beta node streams")
def step_emit_both_messages(context):
alpha_stream = context.graph.stream_router.streams[
f"__{context.graph.name}_node_alpha__"
]
beta_stream = context.graph.stream_router.streams[
f"__{context.graph.name}_node_beta__"
]
alpha_stream.on_next(StreamMessage(content="alpha-payload", metadata={}))
beta_stream.on_next(StreamMessage(content="beta-payload", metadata={}))
@then("each executor should receive its own message independently")
def step_assert_independent_executors(context):
context.alpha_executor.assert_called_once()
context.beta_executor.assert_called_once()
alpha_msg = context.alpha_executor.call_args[0][0]
beta_msg = context.beta_executor.call_args[0][0]
assert isinstance(alpha_msg, StreamMessage)
assert isinstance(beta_msg, StreamMessage)
assert alpha_msg.content == "alpha-payload"
assert beta_msg.content == "beta-payload"
# ── Executor exception handling (M3) ─────────────────────────────────
@given("a langgraph instance with a failing node executor")
def step_prepare_failing_executor(context):
nodes = {"worker": NodeConfig(name="worker", type=NodeType.FUNCTION)}
context.graph = _fresh_graph(GraphConfig(name="graph-fail-exec", nodes=nodes))
context.graph._node_executors["worker"] = MagicMock( # pylint: disable=protected-access
side_effect=RuntimeError("executor boom")
)
context.graph.logger = MagicMock()
@when("a message is emitted on the failing worker stream")
def step_emit_failing_worker_message(context):
worker_stream = context.graph.stream_router.streams[
f"__{context.graph.name}_node_worker__"
]
context.on_next_error = None
try:
worker_stream.on_next(StreamMessage(content="trigger", metadata={}))
except Exception as exc: # pylint: disable=broad-except
context.on_next_error = exc
@then("the exception should be caught and logged without propagating")
def step_assert_exception_caught(context):
assert context.on_next_error is None, (
f"on_next raised unexpectedly: {context.on_next_error}"
)
context.graph.logger.exception.assert_called_once()
args, _kwargs = context.graph.logger.exception.call_args
assert args[0] == "Executor failed for node %s"
assert args[1] == "worker"
# ── Missing observable branch coverage (m3) ──────────────────────────
@given("a langgraph instance with a removed observable")
def step_prepare_removed_observable(context):
nodes = {"worker": NodeConfig(name="worker", type=NodeType.FUNCTION)}
context.graph = _fresh_graph(GraphConfig(name="graph-no-obs", nodes=nodes))
stream_name = f"__{context.graph.name}_node_worker__"
context.graph.stream_router.observables.pop(stream_name, None)
context.graph.logger = MagicMock()
@when("node stream subscriptions are set up")
def step_rerun_subscriptions(context):
# Note: calling _setup_node_stream_subscriptions a second time creates
# duplicate subscriptions for start/end nodes. This is expected and
# harmless in this test context — we only care about the worker branch.
context.graph._setup_node_stream_subscriptions() # pylint: disable=protected-access
@then("no executor should be called and no error should occur")
def step_assert_no_error_on_missing_observable(context):
context.graph.logger.debug.assert_any_call(
"No observable for stream %s; skipping subscription",
f"__{context.graph.name}_node_worker__",
)
context.graph.logger.error.assert_not_called()
context.graph.logger.warning.assert_not_called()
# ── Execute raises when start stream is missing (m1) ─────────────────
@given("a langgraph instance with a removed start stream for execute")
def step_prepare_execute_missing_start(context):
config = GraphConfig(name="graph-exec-no-start")
context.graph = _fresh_graph(config)
context.graph.is_running = True
start_stream = f"__{context.graph.name}_node_start__"
context.graph.stream_router.streams.pop(start_stream, None)
@when("I execute the graph expecting a start stream error")
def step_execute_expecting_error(context):
context.execute_error = None
loop = asyncio.new_event_loop()
asyncio.set_event_loop(loop)
try:
loop.run_until_complete(context.graph.execute({"messages": []}))
except Exception as exc: # pylint: disable=broad-except
context.execute_error = exc
finally:
loop.close()
replacement_loop = asyncio.new_event_loop()
asyncio.set_event_loop(replacement_loop)
context.add_cleanup(replacement_loop.close)
@then("a start stream missing error should be raised from execute")
def step_assert_execute_start_error(context):
assert isinstance(context.execute_error, ValueError)
assert "Start stream not initialized" in str(context.execute_error)
# ── run_coroutine_threadsafe path coverage (M5) ──────────────────────
def _cleanup_bg_loop(context):
"""Shut down the background event loop and join the thread."""
bg_loop = getattr(context, "_bg_loop", None)
bg_thread = getattr(context, "_bg_thread", None)
try:
if bg_loop is not None:
bg_loop.call_soon_threadsafe(bg_loop.stop)
except RuntimeError:
# Loop may already be closed; ensure thread join still happens.
pass
if bg_thread is not None:
bg_thread.join(timeout=5.0)
if bg_loop is not None:
bg_loop.close()
@given(
"a langgraph instance with a scheduler whose loop is running in a background thread"
)
def step_prepare_running_loop_graph(context):
bg_loop = asyncio.new_event_loop()
def run_loop() -> None:
asyncio.set_event_loop(bg_loop)
bg_loop.run_forever()
bg_thread = threading.Thread(target=run_loop, daemon=True)
bg_thread.start()
context._bg_loop = bg_loop
context._bg_thread = bg_thread
# Register cleanup so the background loop is always stopped, even if
# assertions fail in the @then step.
context.add_cleanup(_cleanup_bg_loop, context)
scheduler = AsyncIOScheduler(bg_loop)
nodes = {"worker": NodeConfig(name="worker", type=NodeType.FUNCTION)}
config = GraphConfig(name="graph-running-loop", nodes=nodes)
context.graph = LangGraph(config=config, stream_router=None, scheduler=scheduler)
context.graph.is_running = True
context.graph.state_manager.get_state = MagicMock(return_value=GraphState())
context.graph.state_manager.update_state = MagicMock()
context.graph.nodes["worker"].execute = AsyncMock(return_value={"messages": ["bg"]})
@when("a message is emitted on the worker stream via the running loop path")
def step_emit_via_running_loop(context):
executor = context.graph._node_executors["worker"] # pylint: disable=protected-access
with patch(
"asyncio.run_coroutine_threadsafe", wraps=asyncio.run_coroutine_threadsafe
) as mock_rcts:
context.executor_result = executor(
StreamMessage(content="bg-test", metadata={})
)
context._mock_run_coroutine_threadsafe = mock_rcts
@then("the executor should complete successfully via run_coroutine_threadsafe")
def step_assert_running_loop_executor(context):
assert isinstance(context.executor_result, StreamMessage)
assert context.executor_result.metadata["node"] == "worker"
context.graph.state_manager.update_state.assert_called_once()
# Verify that run_coroutine_threadsafe was actually used (not the
# thread-pool fallback path). The mock retains call history after
# the patch context manager exits in the @when step.
context._mock_run_coroutine_threadsafe.assert_called_once()
# ── Execution history cap test (m5) ──────────────────────────────────
@given("a langgraph instance with execution history at capacity")
def step_prepare_full_history(context):
nodes = {"worker": NodeConfig(name="worker", type=NodeType.FUNCTION)}
config = GraphConfig(name="graph-history-cap", nodes=nodes)
context.graph = _fresh_graph(config)
context.graph.is_running = True
context.graph.state_manager.get_state = MagicMock(return_value=GraphState())
context.graph.state_manager.update_state = MagicMock()
context.graph.nodes["worker"].execute = AsyncMock(return_value={"messages": ["m"]})
# Pre-fill to capacity
for i in range(MAX_EXECUTION_HISTORY):
context.graph.execution_history.append(f"entry-{i}")
context.oldest_entry = "entry-0"
@when("one more executor invocation is recorded")
def step_invoke_one_more(context):
executor = context.graph._node_executors["worker"] # pylint: disable=protected-access
executor(StreamMessage(content="x", metadata={}))
@then(
"the history length should remain at the maximum and the oldest entry should be dropped"
)
def step_assert_history_cap(context):
assert len(context.graph.execution_history) == MAX_EXECUTION_HISTORY
history_list = list(context.graph.execution_history)
assert context.oldest_entry not in history_list
assert history_list[-1] == "worker"
# ── Concurrent state update lock test (m8) ───────────────────────────
@given("a state manager configured for concurrent access")
def step_prepare_concurrent_state_manager(context):
context.state_manager = StateManager(initial_state=GraphState())
context.num_threads = 20
context.updates_per_thread = 10
@when("multiple threads perform state updates simultaneously")
def step_concurrent_updates(context):
barrier = threading.Barrier(context.num_threads)
def updater() -> None:
barrier.wait(timeout=10.0)
for _ in range(context.updates_per_thread):
context.state_manager.update_state(
{"metadata": {"t": threading.current_thread().name}}
)
threads = [threading.Thread(target=updater) for _ in range(context.num_threads)]
for t in threads:
t.start()
for t in threads:
t.join(timeout=30.0)
assert all(not t.is_alive() for t in threads), "Some threads did not complete"
@then("the final execution count should equal the total number of updates")
def step_assert_concurrent_count(context):
expected = context.num_threads * context.updates_per_thread
actual = context.state_manager.get_state().execution_count
assert actual == expected, f"Expected {expected}, got {actual}"
# ── TimeoutError branch coverage (M5) ────────────────────────────────
@given(
"a langgraph instance with a slow executor on a background loop and a short timeout"
)
def step_prepare_slow_executor_bg_loop(context):
bg_loop = asyncio.new_event_loop()
def run_loop() -> None:
asyncio.set_event_loop(bg_loop)
bg_loop.run_forever()
bg_thread = threading.Thread(target=run_loop, daemon=True)
bg_thread.start()
context._bg_loop = bg_loop
context._bg_thread = bg_thread
context.add_cleanup(_cleanup_bg_loop, context)
scheduler = AsyncIOScheduler(bg_loop)
# Use a very short timeout (0.01s) so the test completes quickly.
nodes = {
"worker": NodeConfig(name="worker", type=NodeType.FUNCTION, timeout=0.01),
}
config = GraphConfig(name="graph-timeout-rcts", nodes=nodes)
context.graph = LangGraph(config=config, stream_router=None, scheduler=scheduler)
context.graph.is_running = True
context.graph.state_manager.get_state = MagicMock(return_value=GraphState())
context.graph.state_manager.update_state = MagicMock()
async def slow_execute(state):
await asyncio.sleep(10)
return {"messages": ["never"]}
context.graph.nodes["worker"].execute = slow_execute
@when("the slow executor is invoked via the running loop path")
def step_invoke_slow_executor_rcts(context):
executor = context.graph._node_executors["worker"] # pylint: disable=protected-access
context.timeout_error = None
try:
executor(StreamMessage(content="slow", metadata={}))
except TimeoutError as exc:
context.timeout_error = exc
@then("a TimeoutError should be raised with the node name and timeout duration")
def step_assert_timeout_rcts(context):
assert context.timeout_error is not None, "Expected TimeoutError was not raised"
msg = str(context.timeout_error)
assert "worker" in msg
assert "0.01" in msg
@given("a langgraph instance with a slow executor and a short timeout")
def step_prepare_slow_executor_tp(context):
nodes = {
"worker": NodeConfig(name="worker", type=NodeType.FUNCTION, timeout=0.01),
}
config = GraphConfig(name="graph-timeout-tp", nodes=nodes)
context.graph = _fresh_graph(config)
context.graph.is_running = True
context.graph.state_manager.get_state = MagicMock(return_value=GraphState())
context.graph.state_manager.update_state = MagicMock()
async def slow_execute(state):
await asyncio.sleep(10)
return {"messages": ["never"]}
context.graph.nodes["worker"].execute = slow_execute
@when("the slow executor is invoked via the thread pool path")
def step_invoke_slow_executor_tp(context):
executor = context.graph._node_executors["worker"] # pylint: disable=protected-access
context.timeout_error = None
try:
executor(StreamMessage(content="slow", metadata={}))
except TimeoutError as exc:
context.timeout_error = exc
@then("a TimeoutError should be raised from the thread pool path with the node name")
def step_assert_timeout_tp(context):
assert context.timeout_error is not None, "Expected TimeoutError was not raised"
msg = str(context.timeout_error)
assert "worker" in msg
assert "0.01" in msg
# ── Restart-after-stop test (m7) ─────────────────────────────────────
@given("a langgraph instance that has been started and stopped")
def step_prepare_restart_graph(context):
nodes = {"worker": NodeConfig(name="worker", type=NodeType.FUNCTION)}
config = GraphConfig(name="graph-restart", nodes=nodes)
context.graph = _fresh_graph(config)
context.graph.stream_router.send_message = MagicMock()
context.graph.state_manager.get_state = MagicMock(return_value=GraphState())
context.graph.state_manager.update_state = MagicMock()
context.graph.nodes["worker"].execute = AsyncMock(
return_value={"messages": ["restarted"]}
)
context.graph.start()
context.graph.stop()
assert context.graph.is_running is False
@when("I restart the graph and invoke an executor")
def step_restart_and_invoke(context):
context.graph.start()
# Ensure the pool is cleaned up even if assertions fail later.
context.add_cleanup(context.graph.stop)
# Invoke the executor directly to verify the pool is recreated after
# restart. This does NOT verify the full stream→executor subscription
# path; that is covered by the on_next wiring scenarios above.
executor = context.graph._node_executors["worker"] # pylint: disable=protected-access
context.restart_result = executor(
StreamMessage(content="after-restart", metadata={})
)
@then("the executor should complete successfully after restart")
def step_assert_restart_success(context):
assert isinstance(context.restart_result, StreamMessage)
assert context.restart_result.metadata["node"] == "worker"
assert context.graph.is_running is True
# stop() is handled by context.add_cleanup registered in the @when step.
# ── is_running guard test (m4) ───────────────────────────────────────
@given("a langgraph instance that is not running")
def step_prepare_not_running_graph(context):
nodes = {"worker": NodeConfig(name="worker", type=NodeType.FUNCTION)}
config = GraphConfig(name="graph-not-running", nodes=nodes)
context.graph = _fresh_graph(config)
# Ensure the graph is explicitly not running.
context.graph.is_running = False
@when("I invoke the node executor directly")
def step_invoke_executor_not_running(context):
executor = context.graph._node_executors["worker"] # pylint: disable=protected-access
context.not_running_error = None
try:
executor(StreamMessage(content="should-fail", metadata={}))
except Exception as exc: # pylint: disable=broad-except
context.not_running_error = exc
@then("a RuntimeError should be raised indicating the graph is not running")
def step_assert_not_running_error(context):
assert isinstance(context.not_running_error, RuntimeError)
assert "not running" in str(context.not_running_error)
assert "worker" in str(context.not_running_error)
# ── StateManager.replace_state() direct test (m6) ───────────────────
@given("a state manager with an initial state")
def step_prepare_state_manager_for_replace(context):
context.state_manager = StateManager(initial_state=GraphState())
context.stream_emissions = [] # list[GraphState]
context.state_manager.state_stream.subscribe(
on_next=lambda s: context.stream_emissions.append(s)
)
@when("I call replace_state with a new state")
def step_call_replace_state(context):
context.new_state = GraphState(metadata={"replaced": True}, current_node="new-node")
context.state_manager.replace_state(context.new_state)
@then("get_state should return the new state")
def step_assert_replace_state_get(context):
retrieved = context.state_manager.get_state()
assert retrieved.metadata == {"replaced": True}
assert retrieved.current_node == "new-node"
@then("the state stream should have received a notification")
def step_assert_replace_state_stream(context):
# BehaviorSubject emits the initial value on subscribe, plus the
# replace_state emission, so we expect at least 2 emissions.
assert len(context.stream_emissions) >= 2
@then("the emitted state should be a deep copy not the same object")
def step_assert_replace_state_deep_copy(context):
# The last emission should be a deep copy, not the same object as
# the internal state or the new_state we passed in.
emitted = context.stream_emissions[-1]
assert emitted is not context.new_state
assert emitted.metadata == {"replaced": True}
@then("the emitted state should be distinct from the internal state")
def step_assert_emitted_distinct_from_internal(context):
# After the M1 fix (deep-copy on store), the emitted copy must also
# be a distinct object from the internal state held by the manager.
emitted = context.stream_emissions[-1]
# Direct access intentional: identity check requires the actual
# internal object, not a get_state() copy.
assert emitted is not context.state_manager.state
# ── CancelledError path coverage (M3) ───────────────────────────────
@given("a langgraph instance with a cancellable executor task")
def step_prepare_cancellable_executor(context):
nodes = {
"worker": NodeConfig(name="worker", type=NodeType.FUNCTION),
}
config = GraphConfig(name="graph-cancel", nodes=nodes)
context.graph = _fresh_graph(config)
context.graph.is_running = True
@when("the executor future is cancelled before completion")
def step_cancel_executor_future(context):
context.cancel_error = None
# Patch the executor pool's submit to return a future that raises
# CancelledError on result(). This simulates graph shutdown
# (stop() calls shutdown(wait=True, cancel_futures=True)) cancelling
# a queued-but-not-yet-started task.
cancelled_future: concurrent.futures.Future[StreamMessage] = (
concurrent.futures.Future()
)
# cancel() on a fresh Future is sufficient to put it in CANCELLED state.
cancelled_future.cancel()
executor_fn = context.graph._node_executors["worker"] # pylint: disable=protected-access
with patch.object(
context.graph._executor_pool, # pylint: disable=protected-access
"submit",
return_value=cancelled_future,
):
try:
executor_fn(StreamMessage(content="cancel-me", metadata={}))
except RuntimeError as exc:
context.cancel_error = exc
@then("a RuntimeError should be raised indicating graph stopping")
def step_assert_cancelled_error(context):
assert context.cancel_error is not None, "Expected RuntimeError was not raised"
msg = str(context.cancel_error)
assert "cancelled" in msg and "stopping" in msg, f"Unexpected error message: {msg}"
# ── __del__ safety net test (m3) ─────────────────────────────────────
@given("a langgraph instance that was never stopped")
def step_prepare_unstopped_graph(context):
config = GraphConfig(name="graph-del-test")
context.graph = _fresh_graph(config)
# Ensure the graph has a live executor pool (never stopped).
assert context.graph._executor_pool is not None # pylint: disable=protected-access
@when("__del__ is called on the graph")
def step_call_del(context):
context.del_error = None
try:
context.graph.__del__()
except Exception as exc: # pylint: disable=broad-except
context.del_error = exc
# The @then step "no exception should propagate from __del__" is already
# defined in bridge_coverage_steps.py and is reused here. It checks
# ``context.del_error is None``.
# ── execute() is_running guard test (n1) ─────────────────────────────
@given("a langgraph instance that is not running for execute")
def step_prepare_not_running_for_execute(context):
config = GraphConfig(name="graph-exec-not-running")
context.graph = _fresh_graph(config)
context.graph.is_running = False
@when("I call execute on the non-running graph")
def step_call_execute_not_running(context):
context.execute_not_running_error = None
loop = asyncio.new_event_loop()
asyncio.set_event_loop(loop)
try:
loop.run_until_complete(context.graph.execute({"messages": []}))
except Exception as exc: # pylint: disable=broad-except
context.execute_not_running_error = exc
finally:
loop.close()
replacement_loop = asyncio.new_event_loop()
asyncio.set_event_loop(replacement_loop)
context.add_cleanup(replacement_loop.close)
@then("a RuntimeError should be raised indicating the graph is not running for execute")
def step_assert_execute_not_running_error(context):
assert isinstance(context.execute_not_running_error, RuntimeError)
assert "not running" in str(context.execute_not_running_error)
# ── TimeoutError through stream path (M2 cycle 6) ───────────────────
@given("a langgraph instance with an executor that raises TimeoutError")
def step_prepare_timeout_executor_for_stream(context):
nodes = {"worker": NodeConfig(name="worker", type=NodeType.FUNCTION)}
context.graph = _fresh_graph(GraphConfig(name="graph-timeout-stream", nodes=nodes))
context.graph._node_executors["worker"] = MagicMock( # pylint: disable=protected-access
side_effect=TimeoutError("timed out")
)
context.graph.logger = MagicMock()
@when("a message is emitted on the timeout worker stream")
def step_emit_timeout_worker_stream(context):
worker_stream = context.graph.stream_router.streams[
f"__{context.graph.name}_node_worker__"
]
context.stream_error = None
try:
worker_stream.on_next(StreamMessage(content="trigger-timeout", metadata={}))
except Exception as exc: # pylint: disable=broad-except
context.stream_error = exc
@then("the timeout should be logged at warning level")
def step_assert_timeout_warning_logged(context):
assert context.stream_error is None, (
f"on_next raised unexpectedly: {context.stream_error}"
)
context.graph.logger.warning.assert_called_once()
args, _kwargs = context.graph.logger.warning.call_args
assert args[0] == "Executor timed out for node %s"
assert args[1] == "worker"
@then("logger exception should not be called for timeout")
def step_assert_no_exception_log_for_timeout(context):
context.graph.logger.exception.assert_not_called()
# ── CancelledError in run_coroutine_threadsafe path (M1 cycle 6) ────
@given("a langgraph instance with a scheduler loop and a cancellable coroutine future")
def step_prepare_rcts_cancellable(context):
bg_loop = asyncio.new_event_loop()
def run_loop() -> None:
asyncio.set_event_loop(bg_loop)
bg_loop.run_forever()
bg_thread = threading.Thread(target=run_loop, daemon=True)
bg_thread.start()
context._bg_loop = bg_loop
context._bg_thread = bg_thread
context.add_cleanup(_cleanup_bg_loop, context)
scheduler = AsyncIOScheduler(bg_loop)
nodes = {"worker": NodeConfig(name="worker", type=NodeType.FUNCTION)}
config = GraphConfig(name="graph-cancel-rcts", nodes=nodes)
context.graph = LangGraph(config=config, stream_router=None, scheduler=scheduler)
context.graph.is_running = True
@when("the coroutine future is cancelled before completion")
def step_cancel_coroutine_future(context):
context.cancel_rcts_error = None
cancelled_future: concurrent.futures.Future[StreamMessage] = (
concurrent.futures.Future()
)
# cancel() on a fresh Future is sufficient to put it in CANCELLED state.
cancelled_future.cancel()
executor_fn = context.graph._node_executors["worker"] # pylint: disable=protected-access
with patch(
"asyncio.run_coroutine_threadsafe",
return_value=cancelled_future,
):
try:
executor_fn(StreamMessage(content="cancel-rcts", metadata={}))
except RuntimeError as exc:
context.cancel_rcts_error = exc
@then("a RuntimeError should be raised indicating graph stopping from coroutine path")
def step_assert_rcts_cancelled_error(context):
assert context.cancel_rcts_error is not None, "Expected RuntimeError was not raised"
msg = str(context.cancel_rcts_error)
assert "cancelled" in msg and "stopping" in msg, f"Unexpected error message: {msg}"
+1 -1
View File
@@ -48,7 +48,7 @@ dependencies = [
"tomlkit>=0.13.0", # TOML writing with comment preservation for config CLI
"tenacity>=8.2.0", # Retry framework for service layer resilience
"aiohttp>=3.13.4", # CVE-2026-34515 mitigation: open redirect vulnerability
"a2a-sdk>=0.3.0", # A2A Python SDK — required transport for local (stdio) and server (HTTP) modes (ADR-047)
"a2a-sdk>=0.3.0,<1.0.0", # A2A Python SDK — required transport for local (stdio) and server (HTTP) modes (ADR-047); pinned <1.0.0 (removed legacy A2AClient)
]
[project.optional-dependencies]
+193 -32
View File
@@ -3,11 +3,14 @@
from __future__ import annotations
import asyncio
import collections
import concurrent.futures
import contextlib
import logging
from collections import defaultdict
from collections.abc import Callable
from pathlib import Path
from typing import Any, cast
from typing import Any
from pydantic import BaseModel, Field
from rx.core import Observer # type: ignore[attr-defined]
@@ -24,6 +27,14 @@ from cleveragents.reactive.stream_router import (
StreamType,
)
_OnNextCallable = Callable[[Any], None]
_OnErrorCallable = Callable[[Exception], None]
MAX_EXECUTION_HISTORY: int = 1000
# Default timeout (seconds) for blocking on executor futures.
_DEFAULT_EXECUTOR_TIMEOUT: float = 300.0
class GraphConfig(BaseModel): # pylint: disable=too-many-instance-attributes
name: str
@@ -71,23 +82,39 @@ class LangGraph: # pylint: disable=too-many-instance-attributes
enable_time_travel=config.enable_time_travel,
)
self.execution_history: list[str] = []
self._node_executors: dict[str, Callable[[StreamMessage], StreamMessage]] = {}
# Default: min(32, cpu_count+4) — suitable for I/O-bound node executors
self._executor_pool: concurrent.futures.ThreadPoolExecutor = (
concurrent.futures.ThreadPoolExecutor()
)
self.execution_history: collections.deque[str] = collections.deque(
maxlen=MAX_EXECUTION_HISTORY
)
self.is_running = False
self._create_graph_streams()
self._analyze_graph()
def __del__(self) -> None:
with contextlib.suppress(Exception):
self._executor_pool.shutdown(wait=False)
async def execute(self, input_data: GraphState | dict[str, Any]) -> GraphState:
if not self.is_running:
raise RuntimeError(
f"Cannot execute graph {self.name!r}: graph is not running"
)
state = (
input_data
if isinstance(input_data, GraphState)
else GraphState.from_dict(cast(dict[str, Any], input_data))
else GraphState.from_dict(input_data)
)
# Replace state manager state for a fresh execution context
self.state_manager.state = state
self.state_manager.state_stream.on_next(state)
self.state_manager.replace_state(state)
start_stream = f"__{self.name}_node_start__"
if start_stream in self.stream_router.streams:
self.stream_router.send_message(start_stream, state)
self.stream_router.send_message(start_stream, state.model_copy(deep=True))
else:
raise ValueError("Start stream not initialized")
return self.state_manager.get_state()
def get_execution_history(self) -> list[str]:
@@ -126,13 +153,8 @@ class LangGraph: # pylint: disable=too-many-instance-attributes
def _create_graph_streams(self) -> None:
for node_name in self.nodes:
stream_name = f"__{self.name}_node_{node_name}__"
operators: list[dict[str, Any]] = [
{"type": "map", "params": {"function": f"execute_node_{node_name}"}}
]
self._register_node_executor(node_name)
config = StreamConfig(
name=stream_name, type=StreamType.COLD, operators=operators
)
config = StreamConfig(name=stream_name, type=StreamType.COLD)
self.stream_router.create_stream(config)
self._create_control_streams()
self._setup_node_stream_subscriptions()
@@ -149,25 +171,69 @@ class LangGraph: # pylint: disable=too-many-instance-attributes
self.state_manager.state_stream
)
def _make_on_next_handler(self, name: str) -> _OnNextCallable:
"""Create an ``on_next`` callback for the node named *name*.
The returned closure looks up the executor from
``self._node_executors`` at invocation time, wraps non-StreamMessage
inputs, and catches executor exceptions so they do not propagate
through the RxPy Subject call chain.
"""
def on_next(msg: Any) -> None:
executor = self._node_executors.get(name)
if executor is not None:
if not isinstance(msg, StreamMessage):
msg = StreamMessage(content=msg, metadata={})
try:
# TODO(#10799): route executor return value to successor
# node streams based on self.adjacency_list[name] to
# implement downstream propagation.
executor(msg)
except TimeoutError:
self.logger.warning("Executor timed out for node %s", name)
except RuntimeError as exc:
if "stopping" in str(exc) or "not running" in str(exc):
self.logger.warning(
"Executor interrupted for node %s: %s", name, exc
)
else:
self.logger.exception("Executor failed for node %s", name)
except Exception:
self.logger.exception("Executor failed for node %s", name)
else:
self.logger.warning("No executor found for node %s", name)
return on_next
def _make_on_error_handler(self, stream_name: str) -> _OnErrorCallable:
"""Create an ``on_error`` callback for *stream_name*."""
def on_error(error: Exception) -> None:
self.logger.error("Error in node stream %s: %s", stream_name, error)
return on_error
def _setup_node_stream_subscriptions(self) -> None:
# Intentionally no-op: stream completion requires no cleanup.
def on_completed() -> None:
pass
for node_name in self.nodes:
stream_name = f"__{self.name}_node_{node_name}__"
if stream_name in self.stream_router.observables:
observable = self.stream_router.observables[stream_name]
def on_next(_msg: Any) -> None:
pass
def on_error(error: Exception, name: str = stream_name) -> None:
self.logger.error("Error in node stream %s: %s", name, error)
def on_completed() -> None:
pass
observer = Observer(
on_next=on_next, on_error=on_error, on_completed=on_completed
if stream_name not in self.stream_router.observables:
self.logger.debug(
"No observable for stream %s; skipping subscription",
stream_name,
)
observable.subscribe(observer)
continue
observable = self.stream_router.observables[stream_name]
observer = Observer(
on_next=self._make_on_next_handler(node_name),
on_error=self._make_on_error_handler(stream_name),
on_completed=on_completed,
)
observable.subscribe(observer)
def _register_node_executor(self, node_name: str) -> None:
node = self.nodes[node_name]
@@ -176,6 +242,7 @@ class LangGraph: # pylint: disable=too-many-instance-attributes
state = self.state_manager.get_state()
updates = await node.execute(state)
self.state_manager.update_state(updates, node_id=node_name)
# Thread-safe under CPython GIL; deque.append is atomic.
self.execution_history.append(node_name)
return StreamMessage(
content=updates,
@@ -187,14 +254,87 @@ class LangGraph: # pylint: disable=too-many-instance-attributes
)
def sync_executor(msg: StreamMessage) -> StreamMessage:
def run_async() -> Any:
if not self.is_running:
raise RuntimeError(
f"Cannot execute node {node_name!r}: graph is not running"
)
# Access the scheduler's internal event loop. This relies on
# RxPY's ``AsyncIOScheduler`` exposing ``_loop`` (private API;
# verified: RxPY 4.x). If the attribute is absent the code
# falls back to the thread-pool path gracefully.
scheduler_loop: asyncio.AbstractEventLoop | None = getattr(
self.scheduler, "_loop", None
)
# Determine whether we are already running inside the
# scheduler's event loop thread. Blocking with
# ``future.result()`` from the same thread that owns the loop
# would deadlock, so we only use ``run_coroutine_threadsafe``
# when called from a *different* thread.
try:
running_loop = asyncio.get_running_loop()
except RuntimeError:
running_loop = None
timeout: float = (
node.config.timeout
if node.config.timeout is not None
else _DEFAULT_EXECUTOR_TIMEOUT
)
if (
scheduler_loop is not None
and scheduler_loop.is_running()
and running_loop is not scheduler_loop
):
future = asyncio.run_coroutine_threadsafe(
async_executor(msg), scheduler_loop
)
try:
return future.result(timeout=timeout)
except concurrent.futures.TimeoutError as exc:
# Cancellation is best-effort: a timed-out task may still
# complete and update state if the coroutine has already
# started and does not cooperatively check for cancellation.
future.cancel()
raise TimeoutError(
f"Node {node_name!r} executor timed out after {timeout}s"
) from exc
except concurrent.futures.CancelledError as exc:
raise RuntimeError(
f"Node {node_name!r} execution was cancelled (graph stopping)"
) from exc
def run_async() -> StreamMessage:
return asyncio.run(async_executor(msg))
with concurrent.futures.ThreadPoolExecutor() as executor:
future = executor.submit(run_async)
return future.result()
try:
future_tp = self._executor_pool.submit(run_async)
except RuntimeError as exc:
raise RuntimeError(
f"Cannot execute node {node_name!r}: graph is not running"
) from exc
try:
return future_tp.result(timeout=timeout)
except concurrent.futures.TimeoutError as exc:
# Cancellation is best-effort: ``cancel()`` only succeeds if
# the task has not yet started running. A timed-out task may
# still complete and update state.
if not future_tp.cancel():
self.logger.warning(
"Timed-out task for node %s could not be cancelled; "
"thread slot remains occupied until completion",
node_name,
)
raise TimeoutError(
f"Node {node_name!r} executor timed out after {timeout}s"
) from exc
except concurrent.futures.CancelledError as exc:
raise RuntimeError(
f"Node {node_name!r} execution was cancelled (graph stopping)"
) from exc
setattr(self.stream_router, f"_builtin_execute_node_{node_name}", sync_executor)
self._node_executors[node_name] = sync_executor
def _analyze_graph(self) -> None:
self.adjacency_list: dict[str, list[str]] = defaultdict(list)
@@ -253,6 +393,26 @@ class LangGraph: # pylint: disable=too-many-instance-attributes
raise ValueError("Graph must have start and end nodes")
def start(self, initial_input: Any = None) -> None:
"""Start (or restart) the graph.
Shuts down the previous executor pool before creating a new one
to prevent leaking threads and to ensure in-flight tasks from a
prior cycle cannot corrupt state in the new cycle.
.. warning::
``shutdown(wait=True, cancel_futures=True)`` blocks until all
**already-running** tasks complete. ``cancel_futures`` only
cancels *queued-but-not-started* tasks. If a running task is
stuck in blocking I/O without a timeout, this method will hang
indefinitely. Callers that need bounded startup latency should
ensure all submitted callables respect a timeout or cooperative
cancellation.
"""
# Prevent new submissions while the pool is being swapped.
self.is_running = False
self._executor_pool.shutdown(wait=True, cancel_futures=True)
self._executor_pool = concurrent.futures.ThreadPoolExecutor()
self.is_running = True
start_stream = f"__{self.name}_node_start__"
if start_stream not in self.stream_router.streams:
@@ -263,6 +423,7 @@ class LangGraph: # pylint: disable=too-many-instance-attributes
def stop(self) -> None:
self.is_running = False
self._executor_pool.shutdown(wait=True, cancel_futures=True)
def get_state(self) -> GraphState:
return self.state_manager.get_state()
+124 -47
View File
@@ -4,6 +4,8 @@ from __future__ import annotations
import json
import logging
import os
import threading
from datetime import datetime
from enum import Enum
from pathlib import Path
@@ -91,6 +93,7 @@ class StateManager: # pylint: disable=too-many-instance-attributes
enable_time_travel: bool = False,
):
self.logger = logging.getLogger(__name__)
self._lock = threading.RLock()
self.state = initial_state or GraphState()
self.checkpoint_dir = checkpoint_dir
self.enable_time_travel = enable_time_travel
@@ -104,7 +107,8 @@ class StateManager: # pylint: disable=too-many-instance-attributes
self.checkpoint_dir.mkdir(parents=True, exist_ok=True)
def get_state(self) -> GraphState:
return self.state
with self._lock:
return self.state.model_copy(deep=True)
def update_state(
self,
@@ -112,47 +116,100 @@ class StateManager: # pylint: disable=too-many-instance-attributes
mode: StateUpdateMode = StateUpdateMode.MERGE,
node_id: str | None = None,
) -> GraphState:
if self.is_closed:
raise RuntimeError("StateManager is closed")
if self.enable_time_travel:
snapshot = StateSnapshot(
state=self.state.to_dict(), timestamp=datetime.now(), node_id=node_id
)
self.history.append(snapshot)
if len(self.history) > self.max_history_size:
self.history = self.history[-self.max_history_size :]
"""Apply *updates* to the internal state and notify subscribers.
self.state.update(updates, mode)
self.state.execution_count += 1
self.state_stream.on_next(self.state)
.. note::
self.update_count += 1
if self.checkpoint_dir and self.update_count % self.checkpoint_interval == 0:
self._save_checkpoint()
return self.state
Emission order is **not** guaranteed to match mutation order
under concurrent access. The internal state is always
consistent (mutations are serialized by ``_lock``), but
``state_stream.on_next()`` is called outside the lock, so
Thread B's emission may be observed before Thread A's if the
OS schedules B first after both release the lock.
"""
with self._lock:
if self.is_closed:
raise RuntimeError("StateManager is closed")
if self.enable_time_travel:
snapshot = StateSnapshot(
state=self.state.to_dict(),
timestamp=datetime.now(),
node_id=node_id,
)
self.history.append(snapshot)
if len(self.history) > self.max_history_size:
self.history = self.history[-self.max_history_size :]
def _save_checkpoint(self) -> None:
self.state.update(updates, mode)
self.state.execution_count += 1
self.update_count += 1
checkpoint_data_to_save: dict[str, Any] | None = None
if (
self.checkpoint_dir
and self.update_count % self.checkpoint_interval == 0
):
checkpoint_data_to_save = {
"state": self.state.to_dict(),
"timestamp": datetime.now().strftime("%Y%m%d_%H%M%S"),
"update_count": self.update_count,
}
# Capture a deep copy while still under the lock so subscribers
# receive an immutable point-in-time snapshot.
state_copy = self.state.model_copy(deep=True)
# Perform file I/O outside the lock to avoid blocking other threads.
if checkpoint_data_to_save is not None:
self._save_checkpoint(checkpoint_data_to_save)
# Emit outside the lock to avoid re-entrant deadlock and priority
# inversion when subscribers call back into StateManager.
self.state_stream.on_next(state_copy)
return state_copy
def _save_checkpoint(self, checkpoint_data: dict[str, Any] | None = None) -> None:
"""Persist a checkpoint to disk.
When *checkpoint_data* is ``None`` the method acquires ``_lock``
internally to capture a consistent snapshot. The lock is
reentrant (``threading.RLock``), so calling this method while
``_lock`` is already held by the same thread is safe.
"""
if not self.checkpoint_dir:
return
timestamp = datetime.now().strftime("%Y%m%d_%H%M%S")
checkpoint_file = self.checkpoint_dir / f"checkpoint_{timestamp}.json"
checkpoint_data = {
"state": self.state.to_dict(),
"timestamp": timestamp,
"update_count": self.update_count,
}
if checkpoint_data is None:
# External callers (e.g. bridge) may invoke without pre-built
# data. Acquire the lock to capture a consistent snapshot.
with self._lock:
checkpoint_data = {
"state": self.state.to_dict(),
"timestamp": datetime.now().strftime("%Y%m%d_%H%M%S"),
"update_count": self.update_count,
}
timestamp = checkpoint_data.get(
"timestamp", datetime.now().strftime("%Y%m%d_%H%M%S")
)
update_count = checkpoint_data.get("update_count", 0)
pid = os.getpid()
checkpoint_file = (
self.checkpoint_dir / f"checkpoint_{timestamp}_{update_count}_{pid}.json"
)
checkpoint_file.write_text(
json.dumps(checkpoint_data, indent=2), encoding="utf-8"
)
self.logger.debug("Saved checkpoint: %s", checkpoint_file)
def load_checkpoint(self, checkpoint_file: Path) -> None:
if self.is_closed:
raise RuntimeError("StateManager is closed")
# Read file I/O outside the lock to avoid blocking other threads.
checkpoint_data = json.loads(checkpoint_file.read_text(encoding="utf-8"))
self.state = GraphState.from_dict(checkpoint_data["state"])
self.update_count = checkpoint_data.get("update_count", 0)
self.state_stream.on_next(self.state)
with self._lock:
if self.is_closed:
raise RuntimeError("StateManager is closed")
self.state = GraphState.from_dict(checkpoint_data["state"])
self.update_count = checkpoint_data.get("update_count", 0)
state_copy = self.state.model_copy(deep=True)
self.state_stream.on_next(state_copy)
self.logger.info("Loaded checkpoint: %s", checkpoint_file)
def get_latest_checkpoint(self) -> Path | None:
@@ -164,30 +221,50 @@ class StateManager: # pylint: disable=too-many-instance-attributes
return max(checkpoints, key=lambda p: p.stat().st_mtime)
def time_travel(self, steps_back: int = 1) -> GraphState | None:
if self.is_closed:
raise RuntimeError("StateManager is closed")
if not self.enable_time_travel or not self.history:
return None
if steps_back >= len(self.history):
steps_back = len(self.history) - 1
snapshot = self.history[-(steps_back + 1)]
self.state = GraphState.from_dict(snapshot.state)
self.state_stream.on_next(self.state)
return self.state
with self._lock:
if self.is_closed:
raise RuntimeError("StateManager is closed")
if not self.enable_time_travel or not self.history:
return None
if steps_back >= len(self.history):
steps_back = len(self.history) - 1
snapshot = self.history[-(steps_back + 1)]
self.state = GraphState.from_dict(snapshot.state)
state_copy = self.state.model_copy(deep=True)
self.state_stream.on_next(state_copy)
return state_copy
def replace_state(self, new_state: GraphState) -> None:
"""Atomically replace the internal state and notify subscribers.
.. note::
See :meth:`update_state` for a note on emission ordering
under concurrent access.
"""
with self._lock:
self.state = new_state.model_copy(deep=True)
state_copy = new_state.model_copy(deep=True)
self.state_stream.on_next(state_copy)
def get_state_observable(self) -> Observable:
return self.state_stream
def clear_history(self) -> None:
self.history.clear()
with self._lock:
self.history.clear()
def reset(self, initial_state: GraphState | None = None) -> None:
if self.is_closed:
raise RuntimeError("StateManager is closed")
self.state = initial_state or GraphState()
self.update_count = 0
self.history.clear()
self.state_stream.on_next(self.state)
with self._lock:
if self.is_closed:
raise RuntimeError("StateManager is closed")
self.state = (
initial_state.model_copy(deep=True) if initial_state else GraphState()
)
self.update_count = 0
self.history.clear()
state_copy = self.state.model_copy(deep=True)
self.state_stream.on_next(state_copy)
def close(self) -> None:
"""Mark this manager as closed and complete the state stream.