Kelvin-programmer/rag-chatbot
0
1"""Tests for the RAG engine (uses mocked models for CI speed)."""2 3from unittest.mock import MagicMock, patch4 5import pytest6 7from src.config import Settings8from src.rag_engine import PROMPT_TEMPLATE, RAGEngine9 10 11class TestPromptTemplate:12 def test_has_placeholders(self):13 assert "{context}" in PROMPT_TEMPLATE14 assert "{question}" in PROMPT_TEMPLATE15 16 def test_renders_correctly(self):17 rendered = PROMPT_TEMPLATE.format(context="ctx", question="q")18 assert "ctx" in rendered19 assert "q" in rendered20 21 22class TestRAGEngineQuery:23 @pytest.fixture24 def mock_engine(self, tmp_path):25 settings = Settings(26 embedding_model="all-MiniLM-L12-v2",27 llm_model="google/flan-t5-base",28 vector_store_path=str(tmp_path / "vs"),29 )30 with (31 patch("src.rag_engine.hf_pipeline") as mock_pipe,32 patch("src.rag_engine.VectorStore") as mock_vs_cls,33 ):34 mock_pipe.return_value = MagicMock(35 return_value=[{"generated_text": "Test answer"}]36 )37 mock_vs = MagicMock()38 mock_vs.count = 539 mock_vs.search.return_value = [40 {"text": "Relevant chunk 1", "score": 0.85, "metadata": {"page": 1}},41 {"text": "Relevant chunk 2", "score": 0.72, "metadata": {"page": 2}},42 ]43 mock_vs_cls.return_value = mock_vs44 45 engine = RAGEngine(settings)46 yield engine47 48 def test_query_returns_answer_and_sources(self, mock_engine):49 result = mock_engine.query("What is the policy?")50 assert "answer" in result51 assert "sources" in result52 assert len(result["sources"]) > 053 54 def test_empty_store_returns_fallback(self, tmp_path):55 settings = Settings(56 embedding_model="all-MiniLM-L12-v2",57 llm_model="google/flan-t5-base",58 vector_store_path=str(tmp_path / "vs"),59 )60 with (61 patch("src.rag_engine.hf_pipeline"),62 patch("src.rag_engine.VectorStore") as mock_vs_cls,63 ):64 mock_vs = MagicMock()65 mock_vs.count = 066 mock_vs.search.return_value = []67 mock_vs_cls.return_value = mock_vs68 69 engine = RAGEngine(settings)70 result = engine.query("anything")71 assert "no documents" in result["answer"].lower() or "upload" in result["answer"].lower()72 