from pathlib import Path from types import SimpleNamespace import pytest 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) result1 = await search_tool.function(ctx, "Python") assert "Search limit reached" not in result1 result2 = await search_tool.function(ctx, "JavaScript") assert "Search limit reached" not in result2 @pytest.mark.asyncio async def test_searches_beyond_limit_return_cap_message( self, search_client, search_config ): """Searches beyond max_searches return limit message.""" toolset = create_search_toolset(search_config, max_searches=1) search_tool = toolset.tools["search"] ctx = make_ctx(search_client) result1 = await search_tool.function(ctx, "Python") assert "Search limit reached" not in result1 result2 = await search_tool.function(ctx, "JavaScript") assert "Search limit reached" in result2 @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") result = await search_tool.function(ctx_run1, "Python") assert "Search limit reached" not in result result2 = await search_tool.function(ctx_run1, "JavaScript") assert "Search limit reached" in result2 ctx_run2 = make_ctx(search_client, run_id="run-2") result3 = await search_tool.function(ctx_run2, "Python") assert "Search limit reached" not in result3 @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): result = await search_tool.function(ctx, "Python") assert "Search limit reached" not in result @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