forked from HAL9000/cleveragents-core
187 lines
6.7 KiB
Python
187 lines
6.7 KiB
Python
"""
|
|
Additional BDD steps to cover remaining uncovered branches in state.py.
|
|
"""
|
|
|
|
import tempfile
|
|
from pathlib import Path
|
|
|
|
from behave import given, then, when
|
|
from behave.runner import Context
|
|
|
|
from cleveragents.langgraph.state import GraphState, StateManager, StateUpdateMode
|
|
|
|
|
|
def _ensure_context(context: Context) -> None:
|
|
context.state_instances = getattr(context, "state_instances", {}) or {}
|
|
context.state_managers = getattr(context, "state_managers", {}) or {}
|
|
context.test_data = getattr(context, "test_data", {}) or {}
|
|
context.results = getattr(context, "results", {}) or {}
|
|
|
|
|
|
@given("the state management system is available")
|
|
def step_state_management_available(context: Context):
|
|
_ensure_context(context)
|
|
|
|
|
|
@given("I have a GraphState starting with 49 messages")
|
|
def step_graphstate_with_49_messages(context: Context):
|
|
_ensure_context(context)
|
|
messages = [{"id": i} for i in range(49)]
|
|
context.state_instances["merge_trim"] = GraphState(messages=messages)
|
|
|
|
|
|
@given("I have merge updates containing 5 new messages")
|
|
def step_merge_updates_5_messages(context: Context):
|
|
_ensure_context(context)
|
|
context.test_data["merge_updates"] = {
|
|
"messages": [{"id": 49 + i} for i in range(5)]
|
|
}
|
|
|
|
|
|
@when("I merge the updates into the state")
|
|
def step_merge_updates(context: Context):
|
|
_ensure_context(context)
|
|
state = context.state_instances["merge_trim"]
|
|
state.update(context.test_data["merge_updates"], StateUpdateMode.MERGE)
|
|
|
|
|
|
@then("the message list should be trimmed to 50 messages")
|
|
def step_verify_trimmed_to_50_merge(context: Context):
|
|
_ensure_context(context)
|
|
state = context.state_instances.get("merge_trim")
|
|
assert state is not None
|
|
assert len(state.messages) == 50
|
|
|
|
|
|
@then("the message list should be trimmed to 50 messages after append")
|
|
def step_verify_trimmed_to_50_append(context: Context):
|
|
_ensure_context(context)
|
|
state = context.state_instances.get("append_trim")
|
|
assert state is not None
|
|
assert len(state.messages) == 50
|
|
|
|
|
|
@then("the newest message should be retained after trimming")
|
|
def step_verify_newest_retained_merge(context: Context):
|
|
_ensure_context(context)
|
|
state = context.state_instances["merge_trim"]
|
|
assert state.messages[-1]["id"] == 53
|
|
|
|
|
|
@given("I have a GraphState starting with 50 messages")
|
|
def step_graphstate_with_50_messages(context: Context):
|
|
_ensure_context(context)
|
|
messages = [{"id": i} for i in range(50)]
|
|
context.state_instances["append_trim"] = GraphState(messages=messages)
|
|
|
|
|
|
@given("I have append updates containing 5 new messages")
|
|
def step_append_updates_5_messages(context: Context):
|
|
_ensure_context(context)
|
|
context.test_data["append_updates"] = {
|
|
"messages": [{"id": 50 + i} for i in range(5)]
|
|
}
|
|
|
|
|
|
@when("I append the updates into the state")
|
|
def step_append_updates(context: Context):
|
|
_ensure_context(context)
|
|
state = context.state_instances["append_trim"]
|
|
state.update(context.test_data["append_updates"], StateUpdateMode.APPEND)
|
|
|
|
|
|
@then("the newest appended message should be retained after trimming")
|
|
def step_verify_newest_retained_append(context: Context):
|
|
_ensure_context(context)
|
|
state = context.state_instances["append_trim"]
|
|
assert state.messages[-1]["id"] == 54
|
|
|
|
|
|
@given("I have a StateManager with time travel enabled and history size 2")
|
|
def step_state_manager_time_travel_size_2(context: Context):
|
|
_ensure_context(context)
|
|
manager = StateManager(enable_time_travel=True)
|
|
manager.max_history_size = 2
|
|
context.state_managers["history_trim"] = manager
|
|
|
|
|
|
@when("I perform 4 sequential state updates with time travel")
|
|
def step_perform_4_updates_time_travel(context: Context):
|
|
_ensure_context(context)
|
|
manager = context.state_managers["history_trim"]
|
|
for i in range(4):
|
|
manager.update_state({"execution_count": i}, node_id=f"node_{i}")
|
|
|
|
|
|
@then("only the two most recent history snapshots should remain")
|
|
def step_verify_two_history_snapshots(context: Context):
|
|
manager = context.state_managers["history_trim"]
|
|
assert len(manager.history) == 2
|
|
|
|
|
|
@then("the earliest snapshots should be discarded")
|
|
def step_verify_earliest_discarded(context: Context):
|
|
manager = context.state_managers["history_trim"]
|
|
ids = [snap.node_id for snap in manager.history]
|
|
assert ids == ["node_2", "node_3"]
|
|
|
|
|
|
@given("I have a StateManager with checkpointing enabled and interval 1")
|
|
def step_state_manager_checkpoint_interval_1(context: Context):
|
|
context.test_data = getattr(context, "test_data", {})
|
|
context.state_managers = getattr(context, "state_managers", {})
|
|
checkpoint_dir = Path(tempfile.mkdtemp())
|
|
manager = StateManager(checkpoint_dir=checkpoint_dir)
|
|
manager.checkpoint_interval = 1
|
|
context.state_managers["checkpoint_interval"] = manager
|
|
context.test_data["checkpoint_dir_interval"] = checkpoint_dir
|
|
|
|
|
|
@when("I perform 2 sequential state updates with checkpointing")
|
|
def step_perform_updates_checkpoint(context: Context):
|
|
manager = context.state_managers["checkpoint_interval"]
|
|
for i in range(2):
|
|
manager.update_state({"execution_count": i})
|
|
|
|
|
|
@then("a checkpoint file should exist in the directory")
|
|
def step_verify_checkpoint_exists(context: Context):
|
|
checkpoint_dir = context.test_data["checkpoint_dir_interval"]
|
|
checkpoint_files = list(checkpoint_dir.glob("checkpoint_*.json"))
|
|
assert checkpoint_files, "No checkpoint file found"
|
|
|
|
|
|
@then("checkpoint saving should have been triggered")
|
|
def step_verify_checkpoint_triggered(context: Context):
|
|
manager = context.state_managers["checkpoint_interval"]
|
|
assert manager.update_count >= 2
|
|
|
|
|
|
@given("I have a StateManager with time travel history of 2 snapshots")
|
|
def step_state_manager_history_two(context: Context):
|
|
context.state_managers = getattr(context, "state_managers", {})
|
|
manager = StateManager(enable_time_travel=True)
|
|
manager.update_state({"execution_count": 1, "current_node": "node_a"})
|
|
manager.update_state({"execution_count": 2, "current_node": "node_b"})
|
|
context.state_managers["time_travel_clamp"] = manager
|
|
|
|
|
|
@when("I request time travel 5 steps back")
|
|
def step_request_time_travel_far(context: Context):
|
|
manager = context.state_managers["time_travel_clamp"]
|
|
context.results = getattr(context, "results", {})
|
|
context.results["clamped_state"] = manager.time_travel(5)
|
|
|
|
|
|
@then("time travel should return the earliest available snapshot")
|
|
def step_verify_time_travel_clamped(context: Context):
|
|
result = context.results["clamped_state"]
|
|
assert result is not None
|
|
|
|
|
|
@then("the returned state should reflect the earliest snapshot")
|
|
def step_verify_time_travel_state_values(context: Context):
|
|
result = context.results["clamped_state"]
|
|
assert result.execution_count == 0
|
|
assert result.current_node is None
|