Files
cleveragents-core/tests/unit/langgraph/test_state.py

510 lines
16 KiB
Python

"""
Comprehensive unit tests for langgraph state module.
"""
import json
import tempfile
from datetime import datetime
from pathlib import Path
import pytest
from cleveragents.langgraph.state import (
GraphState,
StateManager,
StateSnapshot,
StateUpdateMode,
)
@pytest.fixture
def state_with_one_message():
"""Fixture for GraphState with one initial message."""
state = GraphState()
state.messages = [{"role": "user", "content": "msg1"}]
return state
class TestStateUpdateMode:
"""Test cases for StateUpdateMode enum."""
def test_state_update_mode_values(self):
"""Test StateUpdateMode enum values."""
assert StateUpdateMode.REPLACE.value == "replace"
assert StateUpdateMode.MERGE.value == "merge"
assert StateUpdateMode.APPEND.value == "append"
class TestStateSnapshot:
"""Test cases for StateSnapshot dataclass."""
def test_state_snapshot_creation(self):
"""Test StateSnapshot creation."""
state_dict = {"messages": [], "metadata": {}}
timestamp = datetime.now()
snapshot = StateSnapshot(state=state_dict, timestamp=timestamp)
assert snapshot.state == state_dict
assert snapshot.timestamp == timestamp
assert snapshot.node_id is None
assert snapshot.metadata == {}
def test_state_snapshot_with_values(self):
"""Test StateSnapshot with all values."""
state_dict = {"messages": [{"role": "user", "content": "test"}]}
timestamp = datetime.now()
snapshot = StateSnapshot(
state=state_dict,
timestamp=timestamp,
node_id="node1",
metadata={"key": "value"}
)
assert snapshot.node_id == "node1"
assert snapshot.metadata["key"] == "value"
class TestGraphState:
"""Test cases for GraphState dataclass."""
def test_graph_state_creation(self):
"""Test GraphState creation with defaults."""
state = GraphState()
assert state.messages == []
assert state.metadata == {}
assert state.current_node is None
assert state.execution_count == 0
assert state.error is None
def test_graph_state_with_values(self):
"""Test GraphState creation with values."""
messages = [{"role": "user", "content": "Hello"}]
metadata = {"key": "value"}
state = GraphState(
messages=messages,
metadata=metadata,
current_node="node1",
execution_count=5,
error="some error"
)
assert state.messages == messages
assert state.metadata == metadata
assert state.current_node == "node1"
assert state.execution_count == 5
assert state.error == "some error"
def test_update_replace_mode(self):
"""Test update with REPLACE mode."""
state = GraphState()
state.current_node = "old_node"
updates = {"current_node": "new_node"}
state.update(updates, StateUpdateMode.REPLACE)
assert state.current_node == "new_node"
def test_update_merge_mode_dict(self):
"""Test update with MERGE mode for dict."""
state = GraphState()
state.metadata = {"key1": "value1"}
updates = {"metadata": {"key2": "value2"}}
state.update(updates, StateUpdateMode.MERGE)
assert state.metadata["key1"] == "value1"
assert state.metadata["key2"] == "value2"
def test_update_merge_mode_list(self, state_with_one_message):
"""Test update with MERGE mode for list."""
state = state_with_one_message
updates = {"messages": [{"role": "assistant", "content": "msg2"}]}
state.update(updates, StateUpdateMode.MERGE)
assert len(state.messages) == 2
assert state.messages[0]["content"] == "msg1"
assert state.messages[1]["content"] == "msg2"
def test_update_merge_mode_non_collection(self):
"""Test update with MERGE mode for non-collection."""
state = GraphState()
state.current_node = "old"
updates = {"current_node": "new"}
state.update(updates, StateUpdateMode.MERGE)
assert state.current_node == "new"
def test_update_append_mode_list_with_list(self, state_with_one_message):
"""Test update with APPEND mode for list appending list."""
state = state_with_one_message
updates = {"messages": [{"role": "assistant", "content": "msg2"}]}
state.update(updates, StateUpdateMode.APPEND)
assert len(state.messages) == 2
def test_update_append_mode_list_with_item(self):
"""Test update with APPEND mode for list appending item."""
state = GraphState()
state.messages = [{"role": "user", "content": "msg1"}]
updates = {"messages": {"role": "assistant", "content": "msg2"}}
state.update(updates, StateUpdateMode.APPEND)
assert len(state.messages) == 2
def test_update_ignores_nonexistent_attributes(self):
"""Test update ignores attributes that don't exist."""
state = GraphState()
updates = {"nonexistent_key": "value"}
state.update(updates, StateUpdateMode.MERGE)
# Should not raise error, just ignore
assert not hasattr(state, "nonexistent_key")
def test_to_dict(self):
"""Test converting state to dict."""
state = GraphState(
messages=[{"role": "user", "content": "test"}],
metadata={"key": "value"},
current_node="node1",
execution_count=3,
error="error"
)
result = state.to_dict()
assert result["messages"] == state.messages
assert result["metadata"] == state.metadata
assert result["current_node"] == "node1"
assert result["execution_count"] == 3
assert result["error"] == "error"
def test_from_dict(self):
"""Test creating state from dict."""
data = {
"messages": [{"role": "user", "content": "test"}],
"metadata": {"key": "value"},
"current_node": "node1",
"execution_count": 5,
"error": None
}
state = GraphState.from_dict(data)
assert state.messages == data["messages"]
assert state.metadata == data["metadata"]
assert state.current_node == "node1"
assert state.execution_count == 5
class TestStateManager:
"""Test cases for StateManager class."""
def test_state_manager_init_defaults(self):
"""Test StateManager initialization with defaults."""
manager = StateManager()
assert manager.state is not None
assert isinstance(manager.state, GraphState)
assert manager.checkpoint_dir is None
assert manager.enable_time_travel is False
assert manager.state_stream is not None
assert manager.history == []
assert manager.update_count == 0
def test_state_manager_init_with_initial_state(self):
"""Test StateManager with initial state."""
initial_state = GraphState(current_node="start")
manager = StateManager(initial_state=initial_state)
assert manager.state.current_node == "start"
def test_state_manager_init_with_checkpoint_dir(self):
"""Test StateManager with checkpoint directory."""
with tempfile.TemporaryDirectory() as tmpdir:
checkpoint_dir = Path(tmpdir) / "checkpoints"
manager = StateManager(checkpoint_dir=checkpoint_dir)
assert manager.checkpoint_dir == checkpoint_dir
assert checkpoint_dir.exists()
def test_state_manager_init_with_time_travel(self):
"""Test StateManager with time travel enabled."""
manager = StateManager(enable_time_travel=True)
assert manager.enable_time_travel is True
def test_get_state(self):
"""Test getting current state."""
initial_state = GraphState(current_node="test")
manager = StateManager(initial_state=initial_state)
state = manager.get_state()
assert state.current_node == "test"
def test_update_state_default_merge(self):
"""Test updating state with default MERGE mode."""
manager = StateManager()
updates = {"current_node": "node1"}
result = manager.update_state(updates)
assert result.current_node == "node1"
assert manager.state.execution_count == 1
assert manager.update_count == 1
def test_update_state_replace_mode(self):
"""Test updating state with REPLACE mode."""
manager = StateManager()
manager.state.metadata = {"old": "value"}
updates = {"metadata": {"new": "value"}}
manager.update_state(updates, mode=StateUpdateMode.REPLACE)
assert "new" in manager.state.metadata
def test_update_state_with_time_travel(self):
"""Test updating state with time travel enabled."""
manager = StateManager(enable_time_travel=True)
updates = {"current_node": "node1"}
manager.update_state(updates, node_id="node1")
assert len(manager.history) == 1
assert manager.history[0].node_id == "node1"
def test_update_state_history_trimming(self):
"""Test that history is trimmed when exceeding max size."""
manager = StateManager(enable_time_travel=True)
manager.max_history_size = 5
# Add more updates than max_history_size
for i in range(10):
manager.update_state({"current_node": f"node{i}"})
assert len(manager.history) == 5
def test_update_state_emits_to_stream(self):
"""Test that state updates emit to stream."""
manager = StateManager()
received_states = []
manager.state_stream.subscribe(received_states.append)
manager.update_state({"current_node": "node1"})
# Should have received the updated state
assert len(received_states) > 0
def test_update_state_triggers_checkpoint(self):
"""Test that state update triggers checkpoint at interval."""
with tempfile.TemporaryDirectory() as tmpdir:
checkpoint_dir = Path(tmpdir)
manager = StateManager(checkpoint_dir=checkpoint_dir)
manager.checkpoint_interval = 2
# Update twice to trigger checkpoint
manager.update_state({"current_node": "node1"})
manager.update_state({"current_node": "node2"})
# Should have created a checkpoint file
checkpoints = list(checkpoint_dir.glob("checkpoint_*.json"))
assert len(checkpoints) == 1
def test_save_checkpoint_without_dir(self):
"""Test save_checkpoint without checkpoint dir does nothing."""
manager = StateManager()
# Should not raise error
manager._save_checkpoint()
def test_save_checkpoint_creates_file(self):
"""Test save_checkpoint creates checkpoint file."""
with tempfile.TemporaryDirectory() as tmpdir:
checkpoint_dir = Path(tmpdir)
manager = StateManager(checkpoint_dir=checkpoint_dir)
manager.state.current_node = "test_node"
manager._save_checkpoint()
checkpoints = list(checkpoint_dir.glob("checkpoint_*.json"))
assert len(checkpoints) == 1
# Verify checkpoint content
with open(checkpoints[0], 'r', encoding='utf-8') as f:
data = json.load(f)
assert data["state"]["current_node"] == "test_node"
def test_load_checkpoint(self):
"""Test loading checkpoint from file."""
with tempfile.TemporaryDirectory() as tmpdir:
checkpoint_dir = Path(tmpdir)
manager = StateManager(checkpoint_dir=checkpoint_dir)
# Create a checkpoint
manager.state.current_node = "original"
manager._save_checkpoint()
# Modify state
manager.state.current_node = "modified"
# Load checkpoint
checkpoint_file = manager.get_latest_checkpoint()
manager.load_checkpoint(checkpoint_file)
assert manager.state.current_node == "original"
def test_get_latest_checkpoint_no_dir(self):
"""Test get_latest_checkpoint with no checkpoint dir."""
manager = StateManager()
result = manager.get_latest_checkpoint()
assert result is None
def test_get_latest_checkpoint_no_files(self):
"""Test get_latest_checkpoint with no checkpoint files."""
with tempfile.TemporaryDirectory() as tmpdir:
checkpoint_dir = Path(tmpdir)
manager = StateManager(checkpoint_dir=checkpoint_dir)
result = manager.get_latest_checkpoint()
assert result is None
def test_get_latest_checkpoint_returns_newest(self):
"""Test get_latest_checkpoint returns newest file."""
with tempfile.TemporaryDirectory() as tmpdir:
checkpoint_dir = Path(tmpdir)
manager = StateManager(checkpoint_dir=checkpoint_dir)
# Create multiple checkpoints
for i in range(3):
manager.state.current_node = f"node{i}"
manager._save_checkpoint()
latest = manager.get_latest_checkpoint()
assert latest is not None
assert isinstance(latest, Path)
with open(latest, 'r', encoding='utf-8') as f:
data = json.load(f)
assert data["state"]["current_node"] == "node2"
def test_time_travel_disabled(self):
"""Test time_travel when disabled returns None."""
manager = StateManager(enable_time_travel=False)
result = manager.time_travel(1)
assert result is None
def test_time_travel_no_history(self):
"""Test time_travel with no history returns None."""
manager = StateManager(enable_time_travel=True)
result = manager.time_travel(1)
assert result is None
def test_time_travel_one_step(self):
"""Test time_travel going back one step."""
manager = StateManager(enable_time_travel=True)
manager.update_state({"current_node": "node1"})
manager.update_state({"current_node": "node2"})
result = manager.time_travel(steps_back=1)
# Time travel goes back in history
# After two updates, history has two snapshots (before each update)
# Going back 1 step means we go to the state before the last update
assert isinstance(result, GraphState)
def test_time_travel_exceeds_history(self):
"""Test time_travel when steps_back exceeds history length."""
manager = StateManager(enable_time_travel=True)
manager.update_state({"current_node": "node1"})
manager.update_state({"current_node": "node2"})
result = manager.time_travel(steps_back=10)
# Should go to the oldest available state
assert isinstance(result, GraphState)
def test_get_state_observable(self):
"""Test getting state observable."""
manager = StateManager()
observable = manager.get_state_observable()
assert observable is not None
assert observable == manager.state_stream
def test_clear_history(self):
"""Test clearing state history."""
manager = StateManager(enable_time_travel=True)
manager.update_state({"current_node": "node1"})
manager.update_state({"current_node": "node2"})
assert len(manager.history) == 2
manager.clear_history()
assert len(manager.history) == 0
def test_reset_to_default(self):
"""Test reset with default state."""
manager = StateManager()
manager.update_state({"current_node": "node1"})
manager.update_state({"current_node": "node2"})
manager.reset()
assert manager.state.current_node is None
assert manager.update_count == 0
assert len(manager.history) == 0
def test_reset_with_initial_state(self):
"""Test reset with custom initial state."""
manager = StateManager()
manager.update_state({"current_node": "node1"})
new_initial = GraphState(current_node="start")
manager.reset(initial_state=new_initial)
assert manager.state.current_node == "start"
assert manager.update_count == 0
def test_reset_emits_to_stream(self):
"""Test that reset emits to stream."""
manager = StateManager()
received_states = []
manager.state_stream.subscribe(on_next=received_states.append)
manager.reset()
# Should have received the reset state
assert len(received_states) > 0
assert isinstance(received_states[0], GraphState)
if __name__ == "__main__":
pytest.main([__file__, "-v"])