forked from cleveragents/cleveragents-core
510 lines
16 KiB
Python
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"])
|
|
|