Rebase onto upstream (a4d95fd)
#12
@@ -0,0 +1,102 @@
|
||||
"""Test media tracking for screenshots."""
|
||||
|
||||
import base64
|
||||
from pathlib import Path
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
import pytest
|
||||
|
||||
from nanobot.agent.loop import AgentLoop
|
||||
from nanobot.agent.tools.anthropic.base import ToolResult
|
||||
from nanobot.bus.events import InboundMessage
|
||||
from nanobot.bus.queue import MessageBus
|
||||
from nanobot.providers.base import LLMResponse, ToolCallRequest
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_provider():
|
||||
"""Create mock LLM provider."""
|
||||
provider = MagicMock()
|
||||
provider.chat = AsyncMock()
|
||||
provider.thinking_budget = 0
|
||||
return provider
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_session_manager():
|
||||
"""Create mock session manager."""
|
||||
session_mgr = MagicMock()
|
||||
session_mgr.load = AsyncMock(return_value={
|
||||
"messages": [],
|
||||
"metadata": {},
|
||||
})
|
||||
session_mgr.save = AsyncMock()
|
||||
return session_mgr
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_bus():
|
||||
"""Create mock message bus."""
|
||||
bus = MagicMock(spec=MessageBus)
|
||||
bus.publish = AsyncMock()
|
||||
return bus
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def agent_loop(mock_provider, mock_session_manager, mock_bus, tmp_path):
|
||||
"""Create agent loop for testing."""
|
||||
return AgentLoop(
|
||||
provider=mock_provider,
|
||||
session_manager=mock_session_manager,
|
||||
bus=mock_bus,
|
||||
workspace=tmp_path,
|
||||
max_iterations=5,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_screenshots_tracked_in_media(agent_loop, mock_provider):
|
||||
"""Test that screenshots from computer tool are tracked and included in OutboundMessage."""
|
||||
# Create a dummy screenshot (1x1 red pixel PNG)
|
||||
dummy_png = base64.b64encode(
|
||||
b'\x89PNG\r\n\x1a\n\x00\x00\x00\rIHDR\x00\x00\x00\x01\x00\x00\x00\x01'
|
||||
b'\x08\x02\x00\x00\x00\x90wS\xde\x00\x00\x00\x0cIDATx\x9cc\xf8\xcf'
|
||||
b'\xc0\x00\x00\x00\x03\x00\x01\x00\x18\xdd\x8d\xb4\x00\x00\x00\x00IEND\xaeB`\x82'
|
||||
).decode()
|
||||
|
||||
# Mock LLM responses
|
||||
mock_provider.chat.side_effect = [
|
||||
# First call: request computer tool
|
||||
LLMResponse(
|
||||
content="Taking screenshot",
|
||||
tool_calls=[ToolCallRequest(id="call_1", name="computer", arguments={"action": "screenshot"})],
|
||||
),
|
||||
# Second call: final response
|
||||
LLMResponse(content="Here's the screenshot"),
|
||||
]
|
||||
|
||||
# Mock the computer tool to return a ToolResult with a screenshot
|
||||
tool_result = ToolResult(
|
||||
output="Screenshot taken",
|
||||
base64_image=dummy_png,
|
||||
)
|
||||
agent_loop.tools.execute = AsyncMock(return_value=tool_result)
|
||||
|
||||
message = InboundMessage(
|
||||
channel="test",
|
||||
chat_id="123",
|
||||
sender_id="user1",
|
||||
content="Take a screenshot",
|
||||
)
|
||||
|
||||
response = await agent_loop._process_message(message)
|
||||
|
||||
# Check that media is included and points to saved file
|
||||
assert response is not None
|
||||
assert response.media is not None
|
||||
assert len(response.media) == 1
|
||||
assert response.media[0].endswith(".png")
|
||||
assert Path(response.media[0]).exists()
|
||||
|
||||
# Cleanup
|
||||
Path(response.media[0]).unlink()
|
||||
Reference in New Issue
Block a user