Skip to content

Commit 2383587

Browse files
committed
test: add checkpointer integration tests for tool-use scenarios
Verify that the checkpointer correctly persists and restores full tool-use conversations (user, assistant with tool call, tool result, and final assistant response). Closes #38
1 parent 5157ff6 commit 2383587

1 file changed

Lines changed: 131 additions & 2 deletions

File tree

tests/agent/test_checkpointer_integration.py

Lines changed: 131 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -1,7 +1,14 @@
1+
from pydantic import BaseModel
2+
13
from cubepi.agent.agent import Agent
4+
from cubepi.agent.types import AgentTool, AgentToolResult
25
from cubepi.checkpointer.memory import MemoryCheckpointer
3-
from cubepi.providers.base import Model
4-
from cubepi.providers.faux import FauxProvider, faux_assistant_message
6+
from cubepi.providers.base import Model, TextContent, ToolResultMessage
7+
from cubepi.providers.faux import (
8+
FauxProvider,
9+
faux_assistant_message,
10+
faux_tool_call,
11+
)
512

613

714
def make_model() -> Model:
@@ -52,3 +59,125 @@ async def test_no_checkpointer_works_as_before(self):
5259
agent = Agent(provider=provider, model=make_model())
5360
await agent.prompt("Hello")
5461
assert len(agent.state.messages) == 2
62+
63+
async def test_tool_use_messages_persisted(self):
64+
"""Checkpointer persists the full tool-use conversation:
65+
user, assistant (tool call), tool result, final assistant."""
66+
67+
class EchoParams(BaseModel):
68+
text: str
69+
70+
async def echo_execute(tool_call_id, params, *, signal=None, on_update=None):
71+
return AgentToolResult(content=[TextContent(text=f"echo: {params.text}")])
72+
73+
echo_tool = AgentTool(
74+
name="echo",
75+
description="Echo the input text",
76+
parameters=EchoParams,
77+
execute=echo_execute,
78+
)
79+
80+
checkpointer = MemoryCheckpointer()
81+
provider = FauxProvider()
82+
# First response: assistant calls the echo tool
83+
# Second response: assistant gives a final text answer
84+
provider.set_responses(
85+
[
86+
faux_assistant_message(
87+
faux_tool_call("echo", {"text": "hello"}, id="tc-1"),
88+
stop_reason="tool_use",
89+
),
90+
faux_assistant_message("Done! The echo said hello."),
91+
]
92+
)
93+
agent = Agent(
94+
provider=provider,
95+
model=make_model(),
96+
tools=[echo_tool],
97+
checkpointer=checkpointer,
98+
thread_id="thread-tool",
99+
)
100+
101+
await agent.prompt("Please echo hello")
102+
103+
data = await checkpointer.load("thread-tool")
104+
assert data is not None
105+
# user + assistant(tool_call) + tool_result + final assistant = 4
106+
assert len(data.messages) == 4
107+
108+
# Verify message roles in order
109+
assert data.messages[0].role == "user"
110+
assert data.messages[1].role == "assistant"
111+
assert data.messages[2].role == "tool_result"
112+
assert data.messages[3].role == "assistant"
113+
114+
# Verify tool result content
115+
tool_result = data.messages[2]
116+
assert isinstance(tool_result, ToolResultMessage)
117+
assert tool_result.tool_call_id == "tc-1"
118+
assert tool_result.tool_name == "echo"
119+
assert any(
120+
hasattr(c, "text") and "echo: hello" in c.text for c in tool_result.content
121+
)
122+
123+
async def test_tool_use_history_restored(self):
124+
"""A second Agent session restores tool-use history and continues."""
125+
126+
class EchoParams(BaseModel):
127+
text: str
128+
129+
async def echo_execute(tool_call_id, params, *, signal=None, on_update=None):
130+
return AgentToolResult(content=[TextContent(text=f"echo: {params.text}")])
131+
132+
echo_tool = AgentTool(
133+
name="echo",
134+
description="Echo the input text",
135+
parameters=EchoParams,
136+
execute=echo_execute,
137+
)
138+
139+
checkpointer = MemoryCheckpointer()
140+
provider = FauxProvider()
141+
142+
# --- First session: tool-use conversation ---
143+
provider.set_responses(
144+
[
145+
faux_assistant_message(
146+
faux_tool_call("echo", {"text": "hi"}, id="tc-1"),
147+
stop_reason="tool_use",
148+
),
149+
faux_assistant_message("The echo returned hi."),
150+
]
151+
)
152+
agent1 = Agent(
153+
provider=provider,
154+
model=make_model(),
155+
tools=[echo_tool],
156+
checkpointer=checkpointer,
157+
thread_id="thread-tool-restore",
158+
)
159+
await agent1.prompt("Echo hi")
160+
# 4 messages after first session
161+
assert len(agent1.state.messages) == 4
162+
163+
# --- Second session: same thread, new Agent ---
164+
provider.set_responses([faux_assistant_message("Sure, continuing.")])
165+
agent2 = Agent(
166+
provider=provider,
167+
model=make_model(),
168+
tools=[echo_tool],
169+
checkpointer=checkpointer,
170+
thread_id="thread-tool-restore",
171+
)
172+
await agent2.prompt("Continue please")
173+
# 4 restored + 2 new (user + assistant) = 6
174+
assert len(agent2.state.messages) == 6
175+
176+
# Verify the restored history kept tool-use messages
177+
assert agent2.state.messages[0].role == "user"
178+
assert agent2.state.messages[1].role == "assistant"
179+
assert agent2.state.messages[2].role == "tool_result"
180+
assert agent2.state.messages[3].role == "assistant"
181+
# New messages
182+
assert agent2.state.messages[4].role == "user"
183+
assert agent2.state.messages[5].role == "assistant"

0 commit comments

Comments
 (0)