"""Tests for marketing_intelligence.agent.anthropic.agent pure functions.""" import asyncio from typing import Any from unittest.mock import AsyncMock, MagicMock, patch from marketing_intelligence.agent.anthropic.agent import ( AnthropicAgent, _key_inputs, _key_result, _prune_old_results, ) class TestKeyInputs: def test_scrape_sound_page(self) -> None: assert _key_inputs("scrape_sound_page", {"sound_id": "s123", "extra": 1}) == { "sound_id": "s123" } def test_build_sound_url(self) -> None: assert _key_inputs("build_sound_url", {"sound_id": "s456"}) == { "sound_id": "s456" } def test_scrape_video(self) -> None: assert _key_inputs("scrape_video", {"url": "https://t.co/v", "x": 1}) == { "url": "https://t.co/v" } def test_scrape_url(self) -> None: assert _key_inputs("scrape_url", {"url": "https://t.co/page"}) == { "url": "https://t.co/page" } def test_scrape_tag_page(self) -> None: assert _key_inputs("scrape_tag_page", {"tag": "dance", "limit": 50}) == { "tag": "dance" } def test_export_results(self) -> None: assert _key_inputs("export_results", {"label": "run1", "other": True}) == { "label": "run1" } def test_close_browser_session(self) -> None: assert _key_inputs("close_browser_session", {"session_id": "sid1"}) == { "session_id": "sid1" } def test_apply_behavior_tactics(self) -> None: assert _key_inputs("apply_behavior_tactics", {"session_id": "sid2"}) == { "session_id": "sid2" } def test_unknown_tool_returns_empty(self) -> None: assert _key_inputs("write_campaign_post", {"post_id": "p1"}) == {} assert _key_inputs("create_browser_session", {"proxy": "gb"}) == {} class TestKeyResult: def test_non_dict_returns_empty(self) -> None: assert _key_result("scrape_video", "string") == {} assert _key_result("scrape_video", 42) == {} assert _key_result("scrape_video", None) == {} assert _key_result("scrape_video", ["a", "b"]) == {} def test_failure_returns_error_with_ok_false(self) -> None: r = _key_result("scrape_video", {"success": False, "error": "timeout"}) assert r == {"ok": False, "error": "timeout"} def test_failure_truncates_long_error(self) -> None: long_err = "e" * 200 r = _key_result("scrape_video", {"success": False, "error": long_err}) assert len(r["error"]) == 120 def test_scrape_tag_page_success(self) -> None: r = _key_result( "scrape_tag_page", {"videos": ["v1", "v2"], "video_count_text": "2.3M"}, ) assert r == {"ok": True, "videos": 2, "video_count_text": "2.3M"} def test_scrape_tag_page_empty_list(self) -> None: r = _key_result("scrape_tag_page", {"videos": [], "video_count_text": None}) assert r["videos"] == 0 def test_scrape_sound_page_success(self) -> None: r = _key_result( "scrape_sound_page", { "title": "Track", "video_urls": ["u1", "u2", "u3"], "video_count_text": "1M", }, ) assert r == { "ok": True, "title": "Track", "video_urls": 3, "video_count_text": "1M", } def test_scrape_video_success(self) -> None: r = _key_result( "scrape_video", {"views": 1000, "likes": 50, "sound_id": "s1"}, ) assert r == {"ok": True, "views": 1000, "likes": 50, "sound_id": "s1"} def test_scrape_url_chars(self) -> None: r = _key_result("scrape_url", {"content": "hello world"}) assert r == {"ok": True, "chars": 11} def test_export_results_success(self) -> None: r = _key_result( "export_results", {"saved_to": "/tmp/out.json", "s3_uri": "s3://bucket/key"}, ) assert r == {"ok": True, "path": "/tmp/out.json", "s3_uri": "s3://bucket/key"} def test_block_signal_non_none_added(self) -> None: r = _key_result("create_browser_session", {"signals": {"block_signal": "soft"}}) assert r.get("block_signal") == "soft" def test_none_block_signal_excluded(self) -> None: r = _key_result("create_browser_session", {"signals": {"block_signal": "none"}}) assert "block_signal" not in r def test_no_signals_key(self) -> None: r = _key_result("create_browser_session", {"session_id": "s1"}) assert "block_signal" not in r assert r["ok"] is True class TestPruneOldResults: def _tool_result_item(self, content: Any) -> dict[str, Any]: return {"type": "tool_result", "tool_use_id": "id1", "content": content} def _user_msg(self, *contents: Any) -> dict[str, Any]: return { "role": "user", "content": [self._tool_result_item(c) for c in contents], } def _all_items(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 item.get("type") == "tool_result" ] def test_prunes_all_but_last_n(self) -> None: msgs = [self._user_msg(str(i)) for i in range(10)] _prune_old_results(msgs, keep_last=3) items = self._all_items(msgs) pruned = [i for i in items if i["content"] == "ok"] kept = [i for i in items if i["content"] != "ok"] assert len(kept) == 3 assert len(pruned) == 7 def test_no_prune_when_fewer_than_keep_last(self) -> None: msgs = [self._user_msg("payload")] _prune_old_results(msgs, keep_last=6) assert msgs[0]["content"][0]["content"] == "payload" 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_items(msgs) assert all(i["content"] != "ok" for i in items) def test_ignores_assistant_messages(self) -> None: msgs: list[dict[str, Any]] = [ {"role": "assistant", "content": [{"type": "text", "text": "hi"}]}, self._user_msg("data"), ] _prune_old_results(msgs, keep_last=6) assert msgs[1]["content"][0]["content"] == "data" def test_ignores_non_tool_result_user_content(self) -> None: msgs: list[dict[str, Any]] = [ { "role": "user", "content": [ {"type": "text", "text": "task"}, self._tool_result_item("payload"), ], } ] _prune_old_results(msgs, keep_last=6) assert msgs[0]["content"][1]["content"] == "payload" def _make_anthropic_agent() -> tuple[AnthropicAgent, MagicMock]: """Return an AnthropicAgent with a mocked Anthropic client.""" mock_client = MagicMock() with patch( "marketing_intelligence.agent.anthropic.agent.anthropic.Anthropic", return_value=mock_client, ): agent = AnthropicAgent(model_id="claude-test") return agent, mock_client def _make_mcp_mock(tool_specs: list | None = None, prompt: str = "system") -> MagicMock: mock_mcp = MagicMock() mock_mcp.list_tools_anthropic = AsyncMock(return_value=tool_specs or []) 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 _make_stream_ctx(stop_reason: str, content: list) -> MagicMock: mock_response = MagicMock() mock_response.stop_reason = stop_reason mock_response.content = content mock_response.usage.input_tokens = 10 mock_response.usage.output_tokens = 5 mock_response.usage.cache_creation_input_tokens = 0 mock_response.usage.cache_read_input_tokens = 0 mock_stream = MagicMock() mock_stream.get_final_message.return_value = mock_response mock_ctx = MagicMock() mock_ctx.__enter__ = MagicMock(return_value=mock_stream) mock_ctx.__exit__ = MagicMock(return_value=False) return mock_ctx class TestAnthropicAgentRun: def test_end_turn_returns_text(self) -> None: agent, mock_client = _make_anthropic_agent() mock_mcp = _make_mcp_mock() text_block = MagicMock() text_block.text = "Analysis complete" ctx = _make_stream_ctx("end_turn", [text_block]) mock_client.messages.stream.return_value = ctx with patch( "marketing_intelligence.agent.anthropic.agent.MCPToolClient", return_value=mock_mcp, ): result = asyncio.run(agent.run("do the task")) assert result == "Analysis complete" def test_end_turn_empty_content_returns_empty_string(self) -> None: agent, mock_client = _make_anthropic_agent() mock_mcp = _make_mcp_mock() ctx = _make_stream_ctx("end_turn", []) mock_client.messages.stream.return_value = ctx with patch( "marketing_intelligence.agent.anthropic.agent.MCPToolClient", return_value=mock_mcp, ): result = asyncio.run(agent.run("task")) assert result == "" def test_unexpected_stop_reason_returns_message(self) -> None: agent, mock_client = _make_anthropic_agent() mock_mcp = _make_mcp_mock() ctx = _make_stream_ctx("max_tokens", []) mock_client.messages.stream.return_value = ctx with patch( "marketing_intelligence.agent.anthropic.agent.MCPToolClient", return_value=mock_mcp, ): result = asyncio.run(agent.run("task")) assert "max_tokens" in result def test_tool_use_then_end_turn_calls_mcp(self) -> None: agent, mock_client = _make_anthropic_agent() mock_mcp = _make_mcp_mock() tool_block = MagicMock() tool_block.type = "tool_use" tool_block.name = "scrape_video" tool_block.id = "tu_1" tool_block.input = {"url": "https://t.co/v"} text_block = MagicMock() text_block.text = "done" ctx1 = _make_stream_ctx("tool_use", [tool_block]) ctx2 = _make_stream_ctx("end_turn", [text_block]) mock_client.messages.stream.side_effect = [ctx1, ctx2] with patch( "marketing_intelligence.agent.anthropic.agent.MCPToolClient", return_value=mock_mcp, ): 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_tool_spec_cache_control_added_to_last(self) -> None: agent, mock_client = _make_anthropic_agent() specs = [ {"name": "tool_a", "description": "a", "input_schema": {}}, {"name": "tool_b", "description": "b", "input_schema": {}}, ] mock_mcp = _make_mcp_mock(tool_specs=specs) text_block = MagicMock() text_block.text = "ok" ctx = _make_stream_ctx("end_turn", [text_block]) mock_client.messages.stream.return_value = ctx captured_tools: list = [] def capture_stream(**kwargs: Any) -> MagicMock: captured_tools.extend(kwargs.get("tools", [])) return ctx mock_client.messages.stream.side_effect = capture_stream with patch( "marketing_intelligence.agent.anthropic.agent.MCPToolClient", return_value=mock_mcp, ): asyncio.run(agent.run("task")) assert "cache_control" in captured_tools[-1] assert "cache_control" not in captured_tools[0] def test_tool_call_exception_returns_error_result(self) -> None: agent, mock_client = _make_anthropic_agent() mock_mcp = _make_mcp_mock() mock_mcp.call_tool = AsyncMock(side_effect=RuntimeError("timeout")) tool_block = MagicMock() tool_block.type = "tool_use" tool_block.name = "scrape_video" tool_block.id = "tu_2" tool_block.input = {} text_block = MagicMock() text_block.text = "recovered" ctx1 = _make_stream_ctx("tool_use", [tool_block]) ctx2 = _make_stream_ctx("end_turn", [text_block]) mock_client.messages.stream.side_effect = [ctx1, ctx2] with patch( "marketing_intelligence.agent.anthropic.agent.MCPToolClient", return_value=mock_mcp, ): result = asyncio.run(agent.run("task")) assert result == "recovered"