RAG_Knowledge_Assistant / tests /test_retriever.py
atara57769's picture
feat: implement modular RAG architecture, query routing, metadata-filtered retrieval, and a comprehensive offline testing suite.
af355c6
Raw
History Blame Contribute Delete
2.84 kB
from unittest.mock import patch, MagicMock
from rag.retriever import retrieve_all
class MockDocument:
def __init__(self, page_content, metadata=None):
self.page_content = page_content
self.metadata = metadata or {}
def test_retrieve_all_no_filter():
"""Verify that retrieval without a filter (or filter='none') delegates to vector store with filter=None."""
with patch('rag.retriever.db') as mock_db:
mock_docs = [
MockDocument("Document 1"),
MockDocument("Document 2")
]
mock_db.similarity_search.return_value = mock_docs
results = retrieve_all("How to buy a toy?", doc_type=None)
assert len(results) == 2
assert results[0].page_content == "Document 1"
mock_db.similarity_search.assert_called_once_with("How to buy a toy?", k=3, filter=None)
def test_retrieve_all_policy_filter():
"""Verify that metadata filter is applied when doc_type is 'policy'."""
with patch('rag.retriever.db') as mock_db, \
patch('rag.retriever.models') as mock_models:
# Configure mock models builder behavior
mock_filter = MagicMock()
mock_models.Filter.return_value = mock_filter
mock_db.similarity_search.return_value = [MockDocument("Policy Doc")]
results = retrieve_all("Return timeline?", doc_type="policy")
assert len(results) == 1
assert results[0].page_content == "Policy Doc"
# Verify models.Filter was called to create the filter
mock_models.Filter.assert_called_once()
mock_db.similarity_search.assert_called_once_with("Return timeline?", k=3, filter=mock_filter)
def test_retrieve_all_product_filter():
"""Verify that metadata filter is applied when doc_type is 'product'."""
with patch('rag.retriever.db') as mock_db, \
patch('rag.retriever.models') as mock_models:
mock_filter = MagicMock()
mock_models.Filter.return_value = mock_filter
mock_db.similarity_search.return_value = [MockDocument("Product Doc")]
results = retrieve_all("Knights castle", doc_type="product")
assert len(results) == 1
assert results[0].page_content == "Product Doc"
mock_models.Filter.assert_called_once()
mock_db.similarity_search.assert_called_once_with("Knights castle", k=3, filter=mock_filter)
def test_retrieve_all_exception_handling():
"""Verify that exceptions in vector store similarity search return an empty list gracefully."""
with patch('rag.retriever.db') as mock_db:
mock_db.similarity_search.side_effect = Exception("Qdrant unavailable")
results = retrieve_all("broken query", doc_type=None)
assert results == []