import pytest
from statebase import StateBase
sb = StateBase(api_key="test-key")
def inject_failure(agent, failure_fn, session_id):
"""Swap the tool that will fail."""
original = agent.tools["search"]
agent.tools["search"] = failure_fn
try:
return agent.run(session_id=session_id, user_input="Do the thing")
finally:
agent.tools["search"] = original
def test_agent_recovers_from_tool_failure():
session = sb.sessions.create(agent_id="test-agent")
def broken_search(query):
raise RuntimeError("search backend down")
result = inject_failure(agent, broken_search, session.id)
# Agent should have caught it and continued
assert "search" in result.lower() or "unavailable" in result.lower()
# State should be intact at a checkpoint
state = sb.sessions.get_state(session_id=session.id)
assert state["stage"] != "errored"
# Turn log shows the failure happened and was handled
turns = sb.sessions.list_turns(session_id=session.id)
assert any("failed" in (t.reasoning or "") for t in turns)
def test_agent_rolls_back_on_state_corruption():
session = sb.sessions.create(
agent_id="test-agent",
initial_state={"important": "value"},
)
sb.sessions.update_state(
session_id=session.id,
state={"important": "value"},
reasoning="checkpoint",
)
# Corrupt
sb.sessions.update_state(
session_id=session.id,
state={"important": None},
reasoning="corruption",
)
# Recover
sb.sessions.rollback(session_id=session.id, version=-2)
state = sb.sessions.get_state(session_id=session.id)
assert state["important"] == "value"