from pathlib import Path from types import SimpleNamespace import pytest from pydantic_ai import ToolFailed from haiku.rag.store.models import SearchResult from haiku.rag.tools.search import create_search_toolset @pytest.fixture(scope="module") def vcr_cassette_dir(): return str(Path(__file__).parent.parent / "cassettes" / "test_search_tools") def make_ctx(client, run_id="test-run"): """Create a lightweight RunContext-like object for direct tool function calls.""" return SimpleNamespace(deps=SimpleNamespace(client=client), run_id=run_id) @pytest.mark.vcr() class TestSearchToolset: """Tests for create_search_toolset.""" def test_create_search_toolset_returns_function_toolset(self, search_config): """create_search_toolset returns a FunctionToolset.""" from pydantic_ai import FunctionToolset toolset = create_search_toolset(search_config) assert isinstance(toolset, FunctionToolset) def test_search_toolset_has_search_tool(self, search_config): """The toolset includes a 'search' tool.""" toolset = create_search_toolset(search_config) # toolset.tools is a dict with tool names as keys assert "search" in toolset.tools @pytest.mark.vcr() class TestSearchToolExecution: """Tests for search tool execution.""" @pytest.mark.asyncio async def test_search_returns_formatted_results(self, search_client, search_config): """Search tool returns formatted results.""" toolset = create_search_toolset(search_config) search_tool = toolset.tools["search"] ctx = make_ctx(search_client) result = await search_tool.function(ctx, "Python") assert "Python" in result or "programming" in result assert "No results found" not in result @pytest.mark.asyncio async def test_search_with_no_results(self, temp_db_path, search_config): """Search tool returns appropriate message when no results.""" from haiku.rag.client import HaikuRAG # Use empty database async with HaikuRAG(temp_db_path, create=True) as empty_client: toolset = create_search_toolset(search_config) search_tool = toolset.tools["search"] ctx = make_ctx(empty_client) result = await search_tool.function(ctx, "anything") assert result == "No results found." @pytest.mark.asyncio async def test_search_with_base_filter(self, search_client, search_config): """Search toolset respects base_filter parameter.""" accumulated: list[SearchResult] = [] toolset = create_search_toolset( search_config, base_filter="title LIKE '%Python%'", on_results=accumulated.extend, ) search_tool = toolset.tools["search"] ctx = make_ctx(search_client) await search_tool.function(ctx, "programming") assert len(accumulated) > 0 for r in accumulated: assert "JavaScript" not in (r.document_title or "") @pytest.mark.asyncio async def test_search_on_results_callback(self, search_client, search_config): """on_results callback receives search results.""" accumulated: list[SearchResult] = [] toolset = create_search_toolset(search_config, on_results=accumulated.extend) search_tool = toolset.tools["search"] ctx = make_ctx(search_client) await search_tool.function(ctx, "Python") assert len(accumulated) > 0 assert any("Python" in r.content for r in accumulated) @pytest.mark.asyncio async def test_search_on_results_accumulates_across_calls( self, search_client, search_config ): """Multiple searches accumulate results via on_results callback.""" accumulated: list[SearchResult] = [] toolset = create_search_toolset(search_config, on_results=accumulated.extend) search_tool = toolset.tools["search"] ctx = make_ctx(search_client) await search_tool.function(ctx, "Python") first_count = len(accumulated) await search_tool.function(ctx, "JavaScript") assert len(accumulated) > first_count @pytest.mark.asyncio async def test_search_without_on_results(self, search_client, search_config): """Search works without on_results callback.""" toolset = create_search_toolset(search_config) search_tool = toolset.tools["search"] ctx = make_ctx(search_client) result = await search_tool.function(ctx, "Python") assert "Python" in result or "programming" in result @pytest.mark.vcr() class TestSearchMaxSearches: """Tests for max_searches cap on search toolset.""" @pytest.mark.asyncio async def test_searches_within_limit_return_results( self, search_client, search_config ): """Searches within max_searches return normal results.""" toolset = create_search_toolset(search_config, max_searches=2) search_tool = toolset.tools["search"] ctx = make_ctx(search_client) assert await search_tool.function(ctx, "Python") assert await search_tool.function(ctx, "JavaScript") @pytest.mark.asyncio async def test_searches_beyond_limit_fail_the_tool( self, search_client, search_config ): """Searches beyond max_searches fail with the limit message.""" toolset = create_search_toolset(search_config, max_searches=1) search_tool = toolset.tools["search"] ctx = make_ctx(search_client) assert await search_tool.function(ctx, "Python") with pytest.raises(ToolFailed, match="Search limit reached"): await search_tool.function(ctx, "JavaScript") @pytest.mark.asyncio async def test_counter_resets_across_runs(self, search_client, search_config): """Search counter resets when run_id changes (new agent run).""" toolset = create_search_toolset(search_config, max_searches=1) search_tool = toolset.tools["search"] ctx_run1 = make_ctx(search_client, run_id="run-1") assert await search_tool.function(ctx_run1, "Python") with pytest.raises(ToolFailed, match="Search limit reached"): await search_tool.function(ctx_run1, "JavaScript") ctx_run2 = make_ctx(search_client, run_id="run-2") assert await search_tool.function(ctx_run2, "Python") @pytest.mark.asyncio async def test_no_limit_by_default(self, search_client, search_config): """Without max_searches, searches are unlimited.""" toolset = create_search_toolset(search_config) search_tool = toolset.tools["search"] ctx = make_ctx(search_client) for _ in range(5): assert await search_tool.function(ctx, "Python") @pytest.fixture async def search_client(temp_db_path): """Create a HaikuRAG client with test data for search tests.""" from haiku.rag.client import HaikuRAG async with HaikuRAG(temp_db_path, create=True) as rag: await rag.create_document( "Python is a programming language. It is widely used for web development.", uri="test://python", title="Python Guide", ) await rag.create_document( "JavaScript runs in the browser. It powers interactive web pages.", uri="test://javascript", title="JavaScript Guide", ) yield rag @pytest.fixture def search_config(): """Default AppConfig for search tests.""" from haiku.rag.config import Config return Config class TestBuildBinaryPartsFromResults: """Picture bytes are attached once per (document, self_ref) pair.""" def test_results_without_image_data_contribute_nothing(self): from haiku.rag.tools.search import build_binary_parts_from_results results = [ SearchResult(content="text only", score=0.5, chunk_id="c1", image_data=None) ] assert build_binary_parts_from_results(results) == [] def test_duplicate_document_and_ref_is_attached_once(self): import base64 from io import BytesIO from PIL import Image as PILImage from haiku.rag.tools.search import build_binary_parts_from_results buf = BytesIO() PILImage.new("RGB", (4, 4), "red").save(buf, format="PNG") png = base64.b64encode(buf.getvalue()).decode() shared = {"#/pictures/0": png} results = [ SearchResult( content="a", score=0.9, chunk_id="c1", document_id="doc-1", image_data=shared, ), SearchResult( content="b", score=0.8, chunk_id="c2", document_id="doc-1", image_data=shared, ), ] parts = build_binary_parts_from_results(results) assert len(parts) == 1