"""Tests for marketing_intelligence.agent.bedrock.agent pure functions.""" import asyncio from typing import Any from unittest.mock import AsyncMock, MagicMock, patch from marketing_intelligence.agent.bedrock.agent import ( BedrockAgent, _key_inputs, _key_result, _prune_old_results, ) class TestKeyInputsBedrock: def test_scrape_tag_page(self) -> None: assert _key_inputs("scrape_tag_page", {"tag": "vibes"}) == {"tag": "vibes"} def test_scrape_video(self) -> None: assert _key_inputs("scrape_video", {"url": "https://t.co/v"}) == { "url": "https://t.co/v" } def test_close_browser_session(self) -> None: assert _key_inputs("close_browser_session", {"session_id": "sid"}) == { "session_id": "sid" } def test_unknown_returns_empty(self) -> None: assert _key_inputs("write_campaign_post", {"post_id": "p1"}) == {} class TestKeyResultBedrock: def test_non_dict_returns_empty(self) -> None: assert _key_result("scrape_video", "str") == {} assert _key_result("scrape_video", None) == {} def test_failure(self) -> None: r = _key_result("scrape_video", {"success": False, "error": "fail"}) assert r == {"ok": False, "error": "fail"} def test_scrape_video_success(self) -> None: r = _key_result("scrape_video", {"views": 500, "likes": 10, "sound_id": "s2"}) assert r == {"ok": True, "views": 500, "likes": 10, "sound_id": "s2"} def test_scrape_tag_page(self) -> None: r = _key_result( "scrape_tag_page", {"videos": ["a", "b", "c"], "video_count_text": "3"} ) assert r["videos"] == 3 def test_block_signal_included(self) -> None: r = _key_result( "create_browser_session", {"signals": {"block_signal": "captcha"}} ) assert r.get("block_signal") == "captcha" class TestPruneOldResultsBedrock: def _user_msg(self, *texts: str) -> dict[str, Any]: return { "role": "user", "content": [ { "toolResult": { "toolUseId": f"id{i}", "content": [{"json": {"value": t}}], "status": "success", } } for i, t in enumerate(texts) ], } def _all_tool_results(self, messages: list[dict[str, Any]]) -> list[dict[str, Any]]: return [ item for msg in messages if msg["role"] == "user" and isinstance(msg["content"], list) for item in msg["content"] if isinstance(item, dict) and "toolResult" in item ] def test_prunes_old_entries(self) -> None: msgs = [self._user_msg(str(i)) for i in range(10)] _prune_old_results(msgs, keep_last=4) items = self._all_tool_results(msgs) pruned = [i for i in items if i["toolResult"]["content"] == [{"text": "ok"}]] kept = [i for i in items if i["toolResult"]["content"] != [{"text": "ok"}]] assert len(pruned) == 6 assert len(kept) == 4 def test_no_prune_when_fewer_than_keep_last(self) -> None: msgs = [self._user_msg("data")] _prune_old_results(msgs, keep_last=6) item = msgs[0]["content"][0] assert item["toolResult"]["content"] == [{"json": {"value": "data"}}] def test_exact_keep_last_unchanged(self) -> None: msgs = [self._user_msg(str(i)) for i in range(6)] _prune_old_results(msgs, keep_last=6) items = self._all_tool_results(msgs) assert all(i["toolResult"]["content"] != [{"text": "ok"}] for i in items) def test_mutates_original_dict(self) -> None: msgs = [self._user_msg("a"), self._user_msg("b"), self._user_msg("c")] _prune_old_results(msgs, keep_last=1) assert msgs[0]["content"][0]["toolResult"]["content"] == [{"text": "ok"}] assert msgs[1]["content"][0]["toolResult"]["content"] == [{"text": "ok"}] def _make_bedrock_agent() -> tuple[BedrockAgent, MagicMock]: mock_client = MagicMock() with patch( "marketing_intelligence.agent.bedrock.agent.boto3.client", return_value=mock_client, ): agent = BedrockAgent(model_id="bedrock-test") return agent, mock_client def _make_bedrock_mcp(prompt: str = "system") -> MagicMock: mock_mcp = MagicMock() mock_mcp.list_tools_bedrock = AsyncMock(return_value=[]) mock_mcp.get_prompt = AsyncMock(return_value=prompt) mock_mcp.call_tool = AsyncMock(return_value={"ok": True}) mock_mcp.__aenter__ = AsyncMock(return_value=mock_mcp) mock_mcp.__aexit__ = AsyncMock(return_value=None) return mock_mcp def _converse_response(stop_reason: str, content: list) -> dict: return { "stopReason": stop_reason, "output": {"message": {"role": "assistant", "content": content}}, } class TestBedrockAgentRun: def test_end_turn_returns_text(self) -> None: agent, _ = _make_bedrock_agent() mock_mcp = _make_bedrock_mcp() resp = _converse_response("end_turn", [{"text": "Bedrock done"}]) with ( patch( "marketing_intelligence.agent.bedrock.agent.MCPToolClient", return_value=mock_mcp, ), patch( "marketing_intelligence.agent.bedrock.agent.asyncio.to_thread", AsyncMock(return_value=resp), ), ): result = asyncio.run(agent.run("task")) assert result == "Bedrock done" def test_end_turn_empty_content_returns_empty_string(self) -> None: agent, _ = _make_bedrock_agent() mock_mcp = _make_bedrock_mcp() resp = _converse_response("end_turn", []) with ( patch( "marketing_intelligence.agent.bedrock.agent.MCPToolClient", return_value=mock_mcp, ), patch( "marketing_intelligence.agent.bedrock.agent.asyncio.to_thread", AsyncMock(return_value=resp), ), ): result = asyncio.run(agent.run("task")) assert result == "" def test_unexpected_stop_reason_returns_message(self) -> None: agent, _ = _make_bedrock_agent() mock_mcp = _make_bedrock_mcp() resp = _converse_response("max_tokens", []) with ( patch( "marketing_intelligence.agent.bedrock.agent.MCPToolClient", return_value=mock_mcp, ), patch( "marketing_intelligence.agent.bedrock.agent.asyncio.to_thread", AsyncMock(return_value=resp), ), ): result = asyncio.run(agent.run("task")) assert "max_tokens" in result def test_tool_use_then_end_turn_calls_mcp(self) -> None: agent, _ = _make_bedrock_agent() mock_mcp = _make_bedrock_mcp() tool_resp = _converse_response( "tool_use", [ { "toolUse": { "toolUseId": "tu_1", "name": "scrape_video", "input": {"url": "https://t.co/v"}, } } ], ) end_resp = _converse_response("end_turn", [{"text": "done"}]) with ( patch( "marketing_intelligence.agent.bedrock.agent.MCPToolClient", return_value=mock_mcp, ), patch( "marketing_intelligence.agent.bedrock.agent.asyncio.to_thread", AsyncMock(side_effect=[tool_resp, end_resp]), ), ): result = asyncio.run(agent.run("task")) assert result == "done" mock_mcp.call_tool.assert_awaited_once_with( "scrape_video", {"url": "https://t.co/v"} ) def test_no_tool_use_blocks_returns_error(self) -> None: agent, _ = _make_bedrock_agent() mock_mcp = _make_bedrock_mcp() resp = _converse_response("tool_use", [{"text": "no tool blocks here"}]) with ( patch( "marketing_intelligence.agent.bedrock.agent.MCPToolClient", return_value=mock_mcp, ), patch( "marketing_intelligence.agent.bedrock.agent.asyncio.to_thread", AsyncMock(return_value=resp), ), ): result = asyncio.run(agent.run("task")) assert "no toolUse blocks" in result def test_tool_call_exception_stored_as_error(self) -> None: agent, _ = _make_bedrock_agent() mock_mcp = _make_bedrock_mcp() mock_mcp.call_tool = AsyncMock(side_effect=RuntimeError("network failure")) tool_resp = _converse_response( "tool_use", [{"toolUse": {"toolUseId": "tu_2", "name": "scrape_video", "input": {}}}], ) end_resp = _converse_response("end_turn", [{"text": "recovered"}]) with ( patch( "marketing_intelligence.agent.bedrock.agent.MCPToolClient", return_value=mock_mcp, ), patch( "marketing_intelligence.agent.bedrock.agent.asyncio.to_thread", AsyncMock(side_effect=[tool_resp, end_resp]), ), ): result = asyncio.run(agent.run("task")) assert result == "recovered"