feat(tests): Add tests for app and cli modules
This commit is contained in:
parent
ef73b53e46
commit
94ff11ad25
2 changed files with 332 additions and 0 deletions
203
tests/test_app.py
Normal file
203
tests/test_app.py
Normal file
|
|
@ -0,0 +1,203 @@
|
|||
import asyncio
|
||||
from pathlib import Path
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
from haiku.rag.app import HaikuRAGApp
|
||||
from haiku.rag.store.models.document import Document
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def app():
|
||||
return HaikuRAGApp(db_path=Path(":memory:"))
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def app():
|
||||
"""Fixture for HaikuRAGApp."""
|
||||
return HaikuRAGApp(db_path=Path(":memory:"))
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_list_documents(app: HaikuRAGApp, monkeypatch):
|
||||
"""Test listing documents."""
|
||||
mock_docs = [
|
||||
Document(id=1, content="doc 1"),
|
||||
Document(id=2, content="doc 2"),
|
||||
]
|
||||
mock_client = AsyncMock()
|
||||
mock_client.list_documents.return_value = mock_docs
|
||||
# The async context manager should return the mock client itself
|
||||
mock_client.__aenter__.return_value = mock_client
|
||||
|
||||
monkeypatch.setattr(app, "_rich_print_document", MagicMock())
|
||||
monkeypatch.setattr(app.console, "print", MagicMock())
|
||||
|
||||
with patch("haiku.rag.app.HaikuRAG", return_value=mock_client):
|
||||
await app.list_documents()
|
||||
|
||||
mock_client.list_documents.assert_called_once()
|
||||
assert app._rich_print_document.call_count == len(mock_docs)
|
||||
app._rich_print_document.assert_any_call(mock_docs[0], truncate=True)
|
||||
app._rich_print_document.assert_any_call(mock_docs[1], truncate=True)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_add_document_from_text(app: HaikuRAGApp, monkeypatch):
|
||||
"""Test adding a document from text."""
|
||||
mock_doc = Document(id=1, content="test document")
|
||||
mock_client = AsyncMock()
|
||||
mock_client.create_document.return_value = mock_doc
|
||||
mock_client.__aenter__.return_value = mock_client
|
||||
|
||||
monkeypatch.setattr(app, "_rich_print_document", MagicMock())
|
||||
mock_print = MagicMock()
|
||||
monkeypatch.setattr(app.console, "print", mock_print)
|
||||
|
||||
with patch("haiku.rag.app.HaikuRAG", return_value=mock_client):
|
||||
await app.add_document_from_text("test document")
|
||||
|
||||
mock_client.create_document.assert_called_once_with("test document")
|
||||
app._rich_print_document.assert_called_once_with(mock_doc, truncate=True)
|
||||
mock_print.assert_called_once_with(
|
||||
"[b]Document with id [cyan]1[/cyan] added successfully.[/b]"
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_add_document_from_source(app: HaikuRAGApp, monkeypatch):
|
||||
"""Test adding a document from a source path."""
|
||||
mock_doc = Document(id=1, content="test document")
|
||||
mock_client = AsyncMock()
|
||||
mock_client.create_document_from_source.return_value = mock_doc
|
||||
mock_client.__aenter__.return_value = mock_client
|
||||
|
||||
monkeypatch.setattr(app, "_rich_print_document", MagicMock())
|
||||
mock_print = MagicMock()
|
||||
monkeypatch.setattr(app.console, "print", mock_print)
|
||||
|
||||
file_path = Path("test.txt")
|
||||
with patch("haiku.rag.app.HaikuRAG", return_value=mock_client):
|
||||
await app.add_document_from_source(file_path)
|
||||
|
||||
mock_client.create_document_from_source.assert_called_once_with(file_path)
|
||||
app._rich_print_document.assert_called_once_with(mock_doc, truncate=True)
|
||||
mock_print.assert_called_once_with(
|
||||
"[b]Document with id [cyan]1[/cyan] added successfully.[/b]"
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_document(app: HaikuRAGApp, monkeypatch):
|
||||
"""Test getting a document."""
|
||||
mock_doc = Document(id=1, content="test document")
|
||||
mock_client = AsyncMock()
|
||||
mock_client.get_document_by_id.return_value = mock_doc
|
||||
mock_client.__aenter__.return_value = mock_client
|
||||
|
||||
monkeypatch.setattr(app, "_rich_print_document", MagicMock())
|
||||
|
||||
with patch("haiku.rag.app.HaikuRAG", return_value=mock_client):
|
||||
await app.get_document(1)
|
||||
|
||||
mock_client.get_document_by_id.assert_called_once_with(1)
|
||||
app._rich_print_document.assert_called_once_with(mock_doc, truncate=False)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_document_not_found(app: HaikuRAGApp, monkeypatch):
|
||||
"""Test getting a document that does not exist."""
|
||||
mock_client = AsyncMock()
|
||||
mock_client.get_document_by_id.return_value = None
|
||||
mock_client.__aenter__.return_value = mock_client
|
||||
|
||||
mock_print = MagicMock()
|
||||
monkeypatch.setattr(app.console, "print", mock_print)
|
||||
|
||||
with patch("haiku.rag.app.HaikuRAG", return_value=mock_client):
|
||||
await app.get_document(1)
|
||||
|
||||
mock_client.get_document_by_id.assert_called_once_with(1)
|
||||
mock_print.assert_called_once_with("[red]Document with id 1 not found.[/red]")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_delete_document(app: HaikuRAGApp, monkeypatch):
|
||||
"""Test deleting a document."""
|
||||
mock_client = AsyncMock()
|
||||
mock_client.__aenter__.return_value = mock_client
|
||||
|
||||
mock_print = MagicMock()
|
||||
monkeypatch.setattr(app.console, "print", mock_print)
|
||||
|
||||
with patch("haiku.rag.app.HaikuRAG", return_value=mock_client):
|
||||
await app.delete_document(1)
|
||||
|
||||
mock_client.delete_document.assert_called_once_with(1)
|
||||
mock_print.assert_called_once_with("[b]Document 1 deleted successfully.[/b]")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_search(app: HaikuRAGApp, monkeypatch):
|
||||
"""Test searching for documents."""
|
||||
mock_results = [("chunk1", 0.9), ("chunk2", 0.8)]
|
||||
mock_client = AsyncMock()
|
||||
mock_client.search.return_value = mock_results
|
||||
mock_client.__aenter__.return_value = mock_client
|
||||
|
||||
monkeypatch.setattr(app, "_rich_print_search_result", MagicMock())
|
||||
|
||||
with patch("haiku.rag.app.HaikuRAG", return_value=mock_client):
|
||||
await app.search("query")
|
||||
|
||||
mock_client.search.assert_called_once_with("query", limit=5, k=60)
|
||||
assert app._rich_print_search_result.call_count == len(mock_results)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_search_no_results(app: HaikuRAGApp, monkeypatch):
|
||||
"""Test searching with no results."""
|
||||
mock_client = AsyncMock()
|
||||
mock_client.search.return_value = []
|
||||
mock_client.__aenter__.return_value = mock_client
|
||||
|
||||
mock_print = MagicMock()
|
||||
monkeypatch.setattr(app.console, "print", mock_print)
|
||||
|
||||
with patch("haiku.rag.app.HaikuRAG", return_value=mock_client):
|
||||
await app.search("query")
|
||||
|
||||
mock_client.search.assert_called_once_with("query", limit=5, k=60)
|
||||
mock_print.assert_called_once_with("[red]No results found.[/red]")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("transport", ["stdio", "sse", "http", None])
|
||||
async def test_serve(app: HaikuRAGApp, monkeypatch, transport):
|
||||
"""Test the serve method with different transports."""
|
||||
mock_server = AsyncMock()
|
||||
mock_watcher = MagicMock()
|
||||
mock_task = asyncio.create_task(asyncio.sleep(0))
|
||||
mock_task.cancel = MagicMock()
|
||||
|
||||
monkeypatch.setattr("haiku.rag.app.create_mcp_server", MagicMock(return_value=mock_server))
|
||||
monkeypatch.setattr("haiku.rag.app.FileWatcher", MagicMock(return_value=mock_watcher))
|
||||
monkeypatch.setattr("asyncio.create_task", MagicMock(return_value=mock_task))
|
||||
|
||||
mock_client = AsyncMock()
|
||||
mock_client.__aenter__.return_value = mock_client
|
||||
|
||||
with patch("haiku.rag.app.HaikuRAG", return_value=mock_client):
|
||||
if transport:
|
||||
await app.serve(transport=transport)
|
||||
else:
|
||||
await app.serve()
|
||||
|
||||
if transport == "stdio":
|
||||
mock_server.run_stdio_async.assert_called_once()
|
||||
elif transport == "sse":
|
||||
mock_server.run_sse_async.assert_called_once_with("sse")
|
||||
else:
|
||||
mock_server.run_http_async.assert_called_once_with("streamable-http")
|
||||
|
||||
mock_task.cancel.assert_called_once()
|
||||
129
tests/test_cli.py
Normal file
129
tests/test_cli.py
Normal file
|
|
@ -0,0 +1,129 @@
|
|||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
from typer.testing import CliRunner
|
||||
|
||||
from haiku.rag.cli import cli
|
||||
|
||||
runner = CliRunner()
|
||||
|
||||
|
||||
def test_list_documents():
|
||||
with patch("haiku.rag.cli.HaikuRAGApp") as mock_app:
|
||||
mock_app_instance = MagicMock()
|
||||
mock_app_instance.list_documents = AsyncMock()
|
||||
mock_app.return_value = mock_app_instance
|
||||
|
||||
result = runner.invoke(cli, ["list"])
|
||||
|
||||
assert result.exit_code == 0
|
||||
mock_app_instance.list_documents.assert_called_once()
|
||||
|
||||
|
||||
def test_add_document_text():
|
||||
with patch("haiku.rag.cli.HaikuRAGApp") as mock_app:
|
||||
mock_app_instance = MagicMock()
|
||||
mock_app_instance.add_document_from_text = AsyncMock()
|
||||
mock_app.return_value = mock_app_instance
|
||||
|
||||
result = runner.invoke(cli, ["add", "test document"])
|
||||
|
||||
assert result.exit_code == 0
|
||||
mock_app_instance.add_document_from_text.assert_called_once_with(
|
||||
text="test document"
|
||||
)
|
||||
|
||||
|
||||
def test_add_document_src():
|
||||
with patch("haiku.rag.cli.HaikuRAGApp") as mock_app:
|
||||
mock_app_instance = MagicMock()
|
||||
mock_app_instance.add_document_from_source = AsyncMock()
|
||||
mock_app.return_value = mock_app_instance
|
||||
|
||||
result = runner.invoke(cli, ["add-src", "test.txt"])
|
||||
|
||||
assert result.exit_code == 0
|
||||
mock_app_instance.add_document_from_source.assert_called_once()
|
||||
|
||||
|
||||
def test_get_document():
|
||||
with patch("haiku.rag.cli.HaikuRAGApp") as mock_app:
|
||||
mock_app_instance = MagicMock()
|
||||
mock_app_instance.get_document = AsyncMock()
|
||||
mock_app.return_value = mock_app_instance
|
||||
|
||||
result = runner.invoke(cli, ["get", "1"])
|
||||
|
||||
assert result.exit_code == 0
|
||||
mock_app_instance.get_document.assert_called_once_with(doc_id=1)
|
||||
|
||||
|
||||
def test_delete_document():
|
||||
with patch("haiku.rag.cli.HaikuRAGApp") as mock_app:
|
||||
mock_app_instance = MagicMock()
|
||||
mock_app_instance.delete_document = AsyncMock()
|
||||
mock_app.return_value = mock_app_instance
|
||||
|
||||
result = runner.invoke(cli, ["delete", "1"])
|
||||
|
||||
assert result.exit_code == 0
|
||||
mock_app_instance.delete_document.assert_called_once_with(doc_id=1)
|
||||
|
||||
|
||||
def test_search():
|
||||
with patch("haiku.rag.cli.HaikuRAGApp") as mock_app:
|
||||
mock_app_instance = MagicMock()
|
||||
mock_app_instance.search = AsyncMock()
|
||||
mock_app.return_value = mock_app_instance
|
||||
|
||||
result = runner.invoke(cli, ["search", "query"])
|
||||
|
||||
assert result.exit_code == 0
|
||||
mock_app_instance.search.assert_called_once_with(query="query", limit=5, k=60)
|
||||
|
||||
|
||||
def test_serve():
|
||||
with patch("haiku.rag.cli.HaikuRAGApp") as mock_app:
|
||||
mock_app_instance = MagicMock()
|
||||
mock_app_instance.serve = AsyncMock()
|
||||
mock_app.return_value = mock_app_instance
|
||||
|
||||
result = runner.invoke(cli, ["serve"])
|
||||
|
||||
assert result.exit_code == 0
|
||||
mock_app_instance.serve.assert_called_once_with(transport=None)
|
||||
|
||||
|
||||
def test_serve_stdio():
|
||||
with patch("haiku.rag.cli.HaikuRAGApp") as mock_app:
|
||||
mock_app_instance = MagicMock()
|
||||
mock_app_instance.serve = AsyncMock()
|
||||
mock_app.return_value = mock_app_instance
|
||||
|
||||
result = runner.invoke(cli, ["serve", "--stdio"])
|
||||
|
||||
assert result.exit_code == 0
|
||||
mock_app_instance.serve.assert_called_once_with(transport="stdio")
|
||||
|
||||
|
||||
def test_serve_sse():
|
||||
with patch("haiku.rag.cli.HaikuRAGApp") as mock_app:
|
||||
mock_app_instance = MagicMock()
|
||||
mock_app_instance.serve = AsyncMock()
|
||||
mock_app.return_value = mock_app_instance
|
||||
|
||||
result = runner.invoke(cli, ["serve", "--sse"])
|
||||
|
||||
assert result.exit_code == 0
|
||||
mock_app_instance.serve.assert_called_once_with(transport="sse")
|
||||
|
||||
|
||||
def test_serve_stdio_and_sse():
|
||||
with patch("haiku.rag.cli.HaikuRAGApp") as mock_app:
|
||||
mock_app_instance = MagicMock()
|
||||
mock_app_instance.serve = AsyncMock()
|
||||
mock_app.return_value = mock_app_instance
|
||||
|
||||
result = runner.invoke(cli, ["serve", "--stdio", "--sse"])
|
||||
|
||||
assert result.exit_code == 1
|
||||
assert "Error: Cannot use both --stdio and --http options" in result.stdout
|
||||
Loading…
Reference in a new issue