from pathlib import Path from unittest.mock import AsyncMock, patch import pytest from typer.testing import CliRunner from haiku.rag.capabilities.rag import RAGState, create_capability from haiku.rag.cli import _cli as cli runner = CliRunner() def test_chat_command(): """Test chat command launches chat TUI.""" with patch("haiku.rag.chat.run_chat") as mock_chat: mock_chat.return_value = None result = runner.invoke(cli, ["chat"]) assert result.exit_code == 0 mock_chat.assert_called_once() def test_run_chat_creates_app_and_runs(temp_db_path: Path): """Test run_chat() eagerly attaches one capability and runs the app.""" with patch("haiku.rag.chat.app.ChatApp") as mock_app: from haiku.rag.chat import run_chat run_chat(db_path=temp_db_path) mock_app.return_value.run.assert_called_once() attached = mock_app.call_args.kwargs["capabilities"] assert len(attached) == 1 assert attached[0].defer_loading is False def test_run_chat_defers_multiple_capabilities(temp_db_path: Path): """Test chat only defers capabilities when routing between multiple choices.""" with patch("haiku.rag.chat.app.ChatApp") as mock_app: from haiku.rag.chat import run_chat run_chat(db_path=temp_db_path, capabilities=["rag", "analysis"]) attached = mock_app.call_args.kwargs["capabilities"] assert len(attached) == 2 assert all(capability.defer_loading for capability in attached) @pytest.mark.parametrize( ("enabled", "expected_model", "expected_vision"), [ (["analysis"], "analysis-model", False), (["rag"], "qa-model", True), (["rag", "analysis"], "qa-model", True), ], ) def test_run_chat_gates_capability_vision_on_driving_model( temp_db_path: Path, enabled, expected_model, expected_vision ): """Analysis-only chat runs on analysis.model; otherwise on qa.model. Every attached capability's vision gate tracks that one driving model.""" from haiku.rag.config.models import AppConfig, ModelConfig config = AppConfig() config.qa.model = ModelConfig(provider="openai", name="qa-model", vision=True) config.analysis.model = ModelConfig( provider="openai", name="analysis-model", vision=False ) captured: dict[str, str] = {} def fake_get_model(model_config, _config): captured["name"] = model_config.name return "resolved-model" with ( patch("haiku.rag.chat.app.ChatApp") as mock_app, patch("haiku.rag.config.get_config", return_value=config), patch("haiku.rag.utils.get_model", side_effect=fake_get_model), ): from haiku.rag.chat import run_chat run_chat(db_path=temp_db_path, capabilities=enabled) assert captured["name"] == expected_model attached = mock_app.call_args.kwargs["capabilities"] assert {capability.vision for capability in attached} == {expected_vision} def _make_mock_client(): """Create a mock HaikuRAG client.""" mock_client = AsyncMock() mock_client.__aenter__ = AsyncMock(return_value=mock_client) mock_client.__aexit__ = AsyncMock(return_value=None) return mock_client def _make_app(db_path: Path, mock_client: AsyncMock | None = None): """Create a ChatApp with mocked HaikuRAG.""" from haiku.rag.chat.app import ChatApp if mock_client is None: mock_client = _make_mock_client() return ChatApp( db_path=db_path, capabilities=[create_capability(db_path=db_path)], read_only=True, ), mock_client def _make_app_with_state(db_path: Path, mock_client: AsyncMock | None = None): """Create a ChatApp with a RAG capability and state.""" from haiku.rag.chat.app import ChatApp if mock_client is None: mock_client = _make_mock_client() return ChatApp( db_path=db_path, capabilities=[create_capability(db_path=db_path)], 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, FlexibleInput.""" from haiku.rag.chat.widgets.chat_history import ChatHistory app, mock_client = _make_app(temp_db_path) with patch("haiku.rag.chat.app.HaikuRAG", return_value=mock_client): async with app.run_test(): chat_history = app.query_one(ChatHistory) assert chat_history is not None from haiku.rag.chat.widgets.prompt import FlexibleInput chat_input = app.query_one(FlexibleInput) assert chat_input is not None @pytest.mark.asyncio async def test_chat_app_quit_binding(temp_db_path: Path): """Test that pressing ctrl+q quits the app.""" app, mock_client = _make_app(temp_db_path) with patch("haiku.rag.chat.app.HaikuRAG", return_value=mock_client): async with app.run_test() as pilot: assert app.is_running await pilot.press("ctrl+q") assert not app.is_running @pytest.mark.asyncio async def test_chat_history_can_add_message(temp_db_path: Path): """Test that ChatHistory can display messages.""" from haiku.rag.chat.widgets.chat_history import ChatHistory app, mock_client = _make_app(temp_db_path) with patch("haiku.rag.chat.app.HaikuRAG", return_value=mock_client): async with app.run_test(): chat_history = app.query_one(ChatHistory) await chat_history.add_message("user", "Hello, how are you?") assert len(chat_history.messages) == 1 assert chat_history.messages[0] == ("user", "Hello, how are you?") await chat_history.add_message("assistant", "I'm doing well, thank you!") assert len(chat_history.messages) == 2 @pytest.mark.asyncio async def test_chat_history_can_add_tool_calls(temp_db_path: Path): """Test that ChatHistory can display inline tool calls.""" from haiku.rag.chat.widgets.chat_history import ChatHistory, ToolCallWidget app, mock_client = _make_app(temp_db_path) with patch("haiku.rag.chat.app.HaikuRAG", return_value=mock_client): async with app.run_test(): chat_history = app.query_one(ChatHistory) tool_widget = await chat_history.add_tool_call( "tool-1", "search", {"query": "test"} ) assert isinstance(tool_widget, ToolCallWidget) assert tool_widget._completed is False chat_history.mark_tool_complete("tool-1") assert tool_widget._completed is True @pytest.mark.asyncio async def test_chat_history_can_add_citations(temp_db_path: Path): """Test that ChatHistory can display inline citations.""" from haiku.rag.chat.widgets.chat_history import ChatHistory, CitationWidget from haiku.rag.store.models.citation import Citation app, mock_client = _make_app(temp_db_path) with patch("haiku.rag.chat.app.HaikuRAG", return_value=mock_client): async with app.run_test(): chat_history = app.query_one(ChatHistory) test_citations = [ Citation( index=1, document_id="doc1", chunk_id="chunk1", document_uri="file:///test/doc1.pdf", document_title="Test Document 1", page_numbers=[1, 2], headings=["Section 1"], content="This is some test content from doc 1", ), Citation( index=2, document_id="doc2", chunk_id="chunk2", document_uri="file:///test/doc2.pdf", document_title="Test Document 2", page_numbers=[5], headings=["Section 2", "Subsection"], content="This is test content from doc 2", ), ] await chat_history.add_citations(test_citations) citation_widgets = chat_history.query(CitationWidget) assert len(list(citation_widgets)) == 2 @pytest.mark.asyncio async def test_chat_history_thinking_indicator(temp_db_path: Path): """Test that ChatHistory can show and hide thinking indicator.""" from haiku.rag.chat.widgets.chat_history import ChatHistory, ThinkingWidget app, mock_client = _make_app(temp_db_path) with patch("haiku.rag.chat.app.HaikuRAG", return_value=mock_client): async with app.run_test() as pilot: chat_history = app.query_one(ChatHistory) await chat_history.show_thinking() thinking = chat_history.query(ThinkingWidget) assert len(list(thinking)) == 1 chat_history.hide_thinking() await pilot.pause() thinking = chat_history.query(ThinkingWidget) assert len(list(thinking)) == 0 @pytest.mark.asyncio async def test_clear_chat_resets_state(temp_db_path: Path): """Test that clearing chat resets state, messages, and conversation id.""" from haiku.rag.chat.widgets.chat_history import ChatHistory app, mock_client = _make_app(temp_db_path) with patch("haiku.rag.chat.app.HaikuRAG", return_value=mock_client): async with app.run_test() as pilot: chat_history = app.query_one(ChatHistory) await chat_history.add_message("user", "Hello") await chat_history.add_message("assistant", "Hi there") assert len(chat_history.messages) == 2 previous_conversation_id = app._conversation_id await app.action_clear_chat() await pilot.pause() assert len(chat_history.messages) == 0 assert app._conversation_id != previous_conversation_id @pytest.mark.asyncio async def test_citation_expand_collapse_with_enter(temp_db_path: Path): """Test that pressing Enter on a focused citation toggles expand/collapse.""" from haiku.rag.chat.widgets.chat_history import ChatHistory, CitationWidget from haiku.rag.store.models.citation import Citation app, mock_client = _make_app(temp_db_path) with patch("haiku.rag.chat.app.HaikuRAG", return_value=mock_client): async with app.run_test() as pilot: chat_history = app.query_one(ChatHistory) test_citation = Citation( index=1, document_id="doc1", chunk_id="chunk1", document_uri="file:///test/doc1.pdf", document_title="Test Document", page_numbers=[1], content="Test content", ) await chat_history.add_citations([test_citation]) citation_widget = chat_history.query_one(CitationWidget) assert citation_widget.collapsed is True citation_widget.focus() await pilot.pause() await pilot.press("enter") await pilot.pause() assert citation_widget.collapsed is False await pilot.press("enter") await pilot.pause() assert citation_widget.collapsed is True @pytest.mark.asyncio async def test_show_citations_renders_from_flat_state(temp_db_path: Path): """Citations in state (flat list[str]) render into the chat history.""" from haiku.rag.chat.app import RAG_STATE_NAMESPACE from haiku.rag.chat.widgets.chat_history import ChatHistory, CitationWidget from haiku.rag.store.models.citation import Citation 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() as pilot: rag_state = RAGState.model_validate(app._state[RAG_STATE_NAMESPACE]) citation = Citation( index=1, document_id="doc1", chunk_id="chunk1", document_uri="file:///test/doc1.pdf", document_title="Test Document", page_numbers=[1], content="Cited content", ) rag_state.citation_index["chunk1"] = citation rag_state.citations.append("chunk1") app._state[RAG_STATE_NAMESPACE] = rag_state.model_dump(mode="json") chat_history = app.query_one(ChatHistory) await app._show_citations_and_programs(chat_history) await pilot.pause() widgets = list(chat_history.query(CitationWidget)) assert len(widgets) == 1 assert widgets[0].citation.chunk_id == "chunk1" @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 = RAGState.model_validate(app._state[RAG_STATE_NAMESPACE]) 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 = RAGState.model_validate(app._state[RAG_STATE_NAMESPACE]) assert rag_state.document_filter is not None # Then clear it app.on_document_filter_modal_filter_changed( DocumentFilterModal.FilterChanged([]) ) rag_state = RAGState.model_validate(app._state[RAG_STATE_NAMESPACE]) assert rag_state.document_filter is None assert app._state["rag"]["document_filter"] is None @pytest.mark.asyncio async def test_chat_app_open_failure_surfaces_real_error(tmp_path: Path): """A failed database open must surface its own error, not an AttributeError from tearing down a client that never opened.""" from haiku.rag.chat.app import ChatApp app = ChatApp(db_path=tmp_path / "missing.lancedb", capabilities=[]) with pytest.raises(FileNotFoundError): async with app.run_test(): pass