|
| 1 | +from pydantic import BaseModel |
| 2 | + |
1 | 3 | from cubepi.agent.agent import Agent |
| 4 | +from cubepi.agent.types import AgentTool, AgentToolResult |
2 | 5 | 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 | +) |
5 | 12 |
|
6 | 13 |
|
7 | 14 | def make_model() -> Model: |
@@ -52,3 +59,125 @@ async def test_no_checkpointer_works_as_before(self): |
52 | 59 | agent = Agent(provider=provider, model=make_model()) |
53 | 60 | await agent.prompt("Hello") |
54 | 61 | 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