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 == []