From c567842420feb0b5b41d61373f184302e7db617c Mon Sep 17 00:00:00 2001 From: Yiorgis Gozadinos Date: Fri, 20 Feb 2026 14:04:05 +0200 Subject: [PATCH] Rewire document filter --- haiku_rag_slim/haiku/rag/chat/app.py | 11 ++++ haiku_rag_slim/haiku/rag/skills/rag.py | 39 ++++++++++--- tests/chat/test_chat_app.py | 76 ++++++++++++++++++++++++++ tests/skills/test_rag.py | 59 +++++++++++++++++++- 4 files changed, 175 insertions(+), 10 deletions(-) diff --git a/haiku_rag_slim/haiku/rag/chat/app.py b/haiku_rag_slim/haiku/rag/chat/app.py index e3bcd107..a5af5f7e 100644 --- a/haiku_rag_slim/haiku/rag/chat/app.py +++ b/haiku_rag_slim/haiku/rag/chat/app.py @@ -8,6 +8,7 @@ from typing import TYPE_CHECKING, Any from haiku.rag.client import HaikuRAG from haiku.rag.config import get_config +from haiku.rag.skills.rag import RAGState from haiku.skills.agent import SkillToolset from haiku.skills.models import Skill @@ -387,4 +388,14 @@ class ChatApp(App): def on_document_filter_modal_filter_changed(self, event: Any) -> None: """Handle document filter changes from modal.""" + from haiku.rag.tools.filters import build_multi_document_filter + self._document_filter = event.selected + + if self._toolset: + rag_state = self._toolset.get_namespace(RAG_STATE_NAMESPACE) + if isinstance(rag_state, RAGState): + rag_state.document_filter = build_multi_document_filter( + self._document_filter + ) + self._state = self._toolset.build_state_snapshot() diff --git a/haiku_rag_slim/haiku/rag/skills/rag.py b/haiku_rag_slim/haiku/rag/skills/rag.py index 296c0520..ca7f9c71 100644 --- a/haiku_rag_slim/haiku/rag/skills/rag.py +++ b/haiku_rag_slim/haiku/rag/skills/rag.py @@ -103,12 +103,22 @@ def create_skill( """ from haiku.rag.client import HaikuRAG + state = ( + ctx.deps.state + if ctx.deps and ctx.deps.state and isinstance(ctx.deps.state, RAGState) + else None + ) + async with HaikuRAG(db_path, config=config, read_only=True) as rag: - results = await rag.search(query, limit=limit) + results = await rag.search( + query, + limit=limit, + filter=state.document_filter if state else None, + ) results = await rag.expand_context(results) - if ctx.deps and ctx.deps.state and isinstance(ctx.deps.state, RAGState): - ctx.deps.state.searches[query] = list(results) + if state: + state.searches[query] = list(results) return "\n\n---\n\n".join( r.format_for_agent(rank=i + 1, total=len(results)) @@ -226,7 +236,10 @@ def create_skill( ) async with HaikuRAG(db_path, config=config, read_only=True) as rag: - answer, citations = await rag.ask(ask_question) + answer, citations = await rag.ask( + ask_question, + filter=state.document_filter if state else None, + ) if ctx.deps and ctx.deps.state and isinstance(ctx.deps.state, RAGState): next_index = len(ctx.deps.state.citations) + 1 @@ -287,18 +300,26 @@ def create_skill( """ from haiku.rag.client import HaikuRAG - async with HaikuRAG(db_path, config=config, read_only=True) as rag: - report = await rag.research(question) + state = ( + ctx.deps.state + if ctx.deps and ctx.deps.state and isinstance(ctx.deps.state, RAGState) + else None + ) - if ctx.deps and ctx.deps.state and isinstance(ctx.deps.state, RAGState): - ctx.deps.state.reports.append( + async with HaikuRAG(db_path, config=config, read_only=True) as rag: + report = await rag.research( + question, filter=state.document_filter if state else None + ) + + if state: + state.reports.append( ResearchEntry( question=question, title=report.title, executive_summary=report.executive_summary, ) ) - ctx.deps.state.qa_history.append( + state.qa_history.append( QAHistoryEntry(question=question, answer=report.executive_summary) ) diff --git a/tests/chat/test_chat_app.py b/tests/chat/test_chat_app.py index 1c41e9b1..d0f38eab 100644 --- a/tests/chat/test_chat_app.py +++ b/tests/chat/test_chat_app.py @@ -5,6 +5,7 @@ import pytest from typer.testing import CliRunner from haiku.rag.cli import _cli as cli +from haiku.rag.skills.rag import RAGState runner = CliRunner() @@ -52,6 +53,29 @@ def _make_app(db_path: Path, mock_client: AsyncMock | None = None): ), mock_client +def _make_app_with_state(db_path: Path, mock_client: AsyncMock | None = None): + """Create a ChatApp with a skill that has RAGState.""" + from haiku.rag.chat.app import ChatApp + from haiku.skills.models import Skill, SkillMetadata, SkillSource + + if mock_client is None: + mock_client = _make_mock_client() + + skill = Skill( + metadata=SkillMetadata(name="rag", description="RAG skill"), + source=SkillSource.ENTRYPOINT, + tools=[], + state_type=RAGState, + state_namespace="rag", + ) + + return ChatApp( + db_path=db_path, + skill=skill, + read_only=True, + ), mock_client + + @pytest.mark.asyncio async def test_chat_app_has_required_widgets(temp_db_path: Path): """Test that ChatApp has the required widgets: ChatHistory, Input.""" @@ -241,3 +265,55 @@ async def test_citation_expand_collapse_with_enter(temp_db_path: Path): await pilot.press("enter") await pilot.pause() assert citation_widget.collapsed is True + + +@pytest.mark.asyncio +async def test_document_filter_updates_rag_state(temp_db_path: Path): + """Test that selecting document filters updates RAGState.document_filter.""" + from haiku.rag.chat.app import RAG_STATE_NAMESPACE + from haiku.rag.chat.widgets.document_filter_modal import DocumentFilterModal + from haiku.rag.tools.filters import build_multi_document_filter + + app, mock_client = _make_app_with_state(temp_db_path) + + with patch("haiku.rag.chat.app.HaikuRAG", return_value=mock_client): + async with app.run_test(): + # Simulate the FilterChanged message + selected = ["AI Overview", "ML Basics"] + app.on_document_filter_modal_filter_changed( + DocumentFilterModal.FilterChanged(selected) + ) + + # RAGState.document_filter should be set + rag_state = app._toolset.get_namespace(RAG_STATE_NAMESPACE) + assert rag_state is not None + expected_filter = build_multi_document_filter(selected) + assert rag_state.document_filter == expected_filter + + # The state snapshot should also reflect the change + assert app._state["rag"]["document_filter"] == expected_filter + + +@pytest.mark.asyncio +async def test_document_filter_cleared_when_empty(temp_db_path: Path): + """Test that clearing all document filters sets document_filter to None.""" + from haiku.rag.chat.app import RAG_STATE_NAMESPACE + from haiku.rag.chat.widgets.document_filter_modal import DocumentFilterModal + + app, mock_client = _make_app_with_state(temp_db_path) + + with patch("haiku.rag.chat.app.HaikuRAG", return_value=mock_client): + async with app.run_test(): + # First set a filter + app.on_document_filter_modal_filter_changed( + DocumentFilterModal.FilterChanged(["AI Overview"]) + ) + rag_state = app._toolset.get_namespace(RAG_STATE_NAMESPACE) + assert rag_state.document_filter is not None + + # Then clear it + app.on_document_filter_modal_filter_changed( + DocumentFilterModal.FilterChanged([]) + ) + assert rag_state.document_filter is None + assert app._state["rag"]["document_filter"] is None diff --git a/tests/skills/test_rag.py b/tests/skills/test_rag.py index fb08576e..deafd57a 100644 --- a/tests/skills/test_rag.py +++ b/tests/skills/test_rag.py @@ -72,6 +72,17 @@ class TestSearchTool: assert len(results) > 0 assert isinstance(results[0], SearchResult) + async def test_search_applies_document_filter_from_state(self, rag_db): + from haiku.rag.skills.rag import RAGState, create_skill + + skill = create_skill(db_path=rag_db) + search = _get_tool(skill, "search") + state = RAGState(document_filter="title = 'AI Overview'") + ctx = _make_ctx(state) + result = await search(ctx, query="artificial intelligence") + assert "AI Overview" in result + assert "ML Basics" not in result + async def test_search_without_state(self, rag_db): from haiku.rag.skills.rag import create_skill @@ -231,7 +242,7 @@ class TestAskTool: call_count = 0 - async def mock_ask(self, question): + async def mock_ask(self, question, **kwargs): nonlocal call_count call_count += 1 if call_count == 1: @@ -252,6 +263,24 @@ class TestAskTool: await ask(ctx, question="Second question") assert state.citations[2].index == 3 + async def test_ask_applies_document_filter_from_state(self, rag_db, monkeypatch): + from haiku.rag.skills.rag import RAGState, create_skill + + captured_kwargs = {} + + async def mock_ask(self, question, **kwargs): + captured_kwargs.update(kwargs) + return ("Answer.", []) + + monkeypatch.setattr(HaikuRAG, "ask", mock_ask) + + skill = create_skill(db_path=rag_db) + ask = _get_tool(skill, "ask") + state = RAGState(document_filter="title = 'AI Overview'") + ctx = _make_ctx(state) + await ask(ctx, question="What is AI?") + assert captured_kwargs.get("filter") == "title = 'AI Overview'" + async def test_ask_includes_prior_qa_context(self, rag_db, monkeypatch): import random @@ -420,6 +449,34 @@ class TestResearchTool: assert state.qa_history[0].question == "What is AI?" assert state.qa_history[0].answer == "AI is transforming industries." + async def test_research_applies_document_filter_from_state( + self, rag_db, monkeypatch + ): + from haiku.rag.skills.rag import RAGState, create_skill + + captured_kwargs = {} + + report = ResearchReport( + title="AI Research", + executive_summary="Summary.", + main_findings=["Finding"], + conclusions=["Conclusion"], + sources_summary="Sources.", + ) + + async def mock_research(self, question, **kwargs): + captured_kwargs.update(kwargs) + return report + + monkeypatch.setattr(HaikuRAG, "research", mock_research) + + skill = create_skill(db_path=rag_db) + research = _get_tool(skill, "research") + state = RAGState(document_filter="title = 'AI Overview'") + ctx = _make_ctx(state) + await research(ctx, question="What is AI?") + assert captured_kwargs.get("filter") == "title = 'AI Overview'" + async def test_research_without_state(self, rag_db, monkeypatch): from haiku.rag.skills.rag import create_skill