| """ |
| Unit tests for vector_store module |
| Tests ChromaDB vector store operations |
| """ |
|
|
| import unittest |
| from unittest.mock import MagicMock, Mock, patch |
|
|
| from src.vector_store import (_calculate_similarity_impl, |
| _generate_embeddings_impl, _process_context_impl, |
| calculate_similarity, generate_embeddings, |
| process_context) |
|
|
|
|
| class TestVectorStore(unittest.TestCase): |
| """Test cases for vector_store module""" |
|
|
| def setUp(self): |
| """Set up test fixtures""" |
| |
| self.mock_doc = Mock() |
| self.mock_doc.page_content = "Test document content" |
| self.mock_doc.metadata = { |
| "id": "KB001", |
| "question": "Test question?", |
| "content": "Test answer.", |
| "section": "Test", |
| } |
|
|
| self.mock_documents = [self.mock_doc] |
|
|
| @patch("src.vector_store.get_embedding_model") |
| def test_generate_embeddings_impl(self, mock_get_embedding_model): |
| """Test internal embedding generation implementation""" |
| mock_model = Mock() |
| mock_model.encode.side_effect = [ |
| [0.1, 0.2, 0.3], |
| [[0.2, 0.3, 0.4]], |
| ] |
| mock_get_embedding_model.return_value = mock_model |
|
|
| query = "Test query" |
| query_emb, doc_embs = _generate_embeddings_impl(query, self.mock_documents) |
|
|
| |
| self.assertEqual(mock_model.encode.call_count, 2) |
|
|
| |
| self.assertEqual(query_emb, [0.1, 0.2, 0.3]) |
| self.assertEqual(len(doc_embs), 1) |
| self.assertEqual(doc_embs[0], [0.2, 0.3, 0.4]) |
|
|
| @patch("src.vector_store.get_embedding_model") |
| def test_generate_embeddings_with_timer(self, mock_get_embedding_model): |
| """Test embedding generation with timer""" |
| |
| mock_model = Mock() |
| mock_model.encode.side_effect = [ |
| [0.1, 0.2, 0.3], |
| [[0.1, 0.2, 0.3]], |
| ] |
| mock_get_embedding_model.return_value = mock_model |
|
|
| mock_timer = Mock() |
| mock_timer.time_step = MagicMock() |
| mock_timer.time_step.return_value.__enter__ = Mock() |
| mock_timer.time_step.return_value.__exit__ = Mock() |
|
|
| generate_embeddings("Test", self.mock_documents, timer=mock_timer) |
|
|
| |
| mock_timer.time_step.assert_called_once_with("embedding_generation") |
|
|
| @patch("src.vector_store.get_embedding_model") |
| def test_generate_embeddings_multiple_docs(self, mock_get_embedding_model): |
| """Test embedding generation with multiple documents""" |
| |
| mock_doc2 = Mock() |
| mock_doc2.page_content = "Second document" |
| docs = [self.mock_doc, mock_doc2] |
|
|
| |
| mock_model = Mock() |
| mock_model.encode.side_effect = [ |
| [0.1, 0.2, 0.3], |
| [[0.2, 0.3, 0.4], [0.3, 0.4, 0.5]], |
| ] |
| mock_get_embedding_model.return_value = mock_model |
|
|
| query_emb, doc_embs = _generate_embeddings_impl("Test", docs) |
|
|
| |
| self.assertEqual(len(doc_embs), 2) |
| self.assertEqual(mock_model.encode.call_count, 2) |
|
|
| @patch("src.vector_store.get_embedding_model") |
| def test_generate_embeddings_empty_documents(self, mock_get_embedding_model): |
| """Test embedding generation when no documents are retrieved""" |
| mock_model = Mock() |
| mock_model.encode.return_value = [0.1, 0.2, 0.3] |
| mock_get_embedding_model.return_value = mock_model |
|
|
| query_emb, doc_embs = _generate_embeddings_impl("Test", []) |
|
|
| self.assertEqual(query_emb, [0.1, 0.2, 0.3]) |
| self.assertEqual(doc_embs, []) |
| mock_model.encode.assert_called_once_with("Test") |
|
|
| def test_calculate_similarity_impl(self): |
| """Test internal similarity calculation implementation""" |
| query_embedding = [1.0, 0.0, 0.0] |
| doc_embeddings = [ |
| [1.0, 0.0, 0.0], |
| [0.0, 1.0, 0.0], |
| [0.5, 0.5, 0.0], |
| ] |
|
|
| scores = _calculate_similarity_impl(query_embedding, doc_embeddings) |
|
|
| |
| self.assertEqual(len(scores), 3) |
| self.assertAlmostEqual(scores[0], 1.0, places=5) |
| self.assertAlmostEqual(scores[1], 0.0, places=5) |
| self.assertGreater(scores[2], 0.0) |
| self.assertLess(scores[2], 1.0) |
|
|
| def test_calculate_similarity_with_timer(self): |
| """Test similarity calculation with timer""" |
| mock_timer = Mock() |
| mock_timer.time_step = MagicMock() |
| mock_timer.time_step.return_value.__enter__ = Mock() |
| mock_timer.time_step.return_value.__exit__ = Mock() |
|
|
| query_emb = [1.0, 0.0, 0.0] |
| doc_embs = [[1.0, 0.0, 0.0]] |
|
|
| calculate_similarity(query_emb, doc_embs, timer=mock_timer) |
|
|
| |
| mock_timer.time_step.assert_called_once_with("similarity_calculation") |
|
|
| def test_process_context_impl(self): |
| """Test internal context processing implementation""" |
| |
| results = [] |
| for i in range(3): |
| mock_result = Mock() |
| mock_result.metadata = { |
| "id": f"KB00{i+1}", |
| "question": f"Question {i+1}?", |
| "content": f"Answer {i+1}.", |
| } |
| results.append(mock_result) |
|
|
| |
| cosine_scores = [0.7, 0.5, 0.9] |
|
|
| context, source_ids, knowledge_pairs = _process_context_impl( |
| results, cosine_scores, max_results=2 |
| ) |
|
|
| |
| self.assertEqual(len(source_ids), 2) |
| self.assertEqual(len(knowledge_pairs), 2) |
|
|
| |
| self.assertEqual(source_ids[0], "KB003") |
| self.assertEqual(knowledge_pairs[0][0], "Question 3?") |
|
|
| |
| self.assertIn("Knowledge Entry 1:", context) |
| self.assertIn("Knowledge Entry 2:", context) |
| self.assertIn("Q: Question 3?", context) |
| self.assertIn("A: Answer 3.", context) |
|
|
| def test_process_context_with_timer(self): |
| """Test context processing with timer""" |
| mock_result = Mock() |
| mock_result.metadata = {"id": "KB001", "question": "Q?", "content": "A."} |
|
|
| mock_timer = Mock() |
| mock_timer.time_step = MagicMock() |
| mock_timer.time_step.return_value.__enter__ = Mock() |
| mock_timer.time_step.return_value.__exit__ = Mock() |
|
|
| process_context([mock_result], [0.9], timer=mock_timer) |
|
|
| |
| mock_timer.time_step.assert_called_once_with("context_processing") |
|
|
| def test_process_context_max_results(self): |
| """Test that max_results parameter limits output""" |
| |
| results = [] |
| for i in range(5): |
| mock_result = Mock() |
| mock_result.metadata = { |
| "id": f"KB00{i}", |
| "question": f"Q{i}?", |
| "content": f"A{i}.", |
| } |
| results.append(mock_result) |
|
|
| scores = [0.9, 0.8, 0.7, 0.6, 0.5] |
|
|
| |
| context, source_ids, knowledge_pairs = _process_context_impl( |
| results, scores, max_results=3 |
| ) |
|
|
| |
| self.assertEqual(len(source_ids), 3) |
| self.assertEqual(len(knowledge_pairs), 3) |
|
|
| def test_process_context_formatting(self): |
| """Test context formatting details""" |
| mock_result = Mock() |
| mock_result.metadata = { |
| "id": "KB001", |
| "question": "Test question?", |
| "content": "Test answer.", |
| } |
|
|
| context, _, _ = _process_context_impl([mock_result], [0.9], max_results=1) |
|
|
| |
| self.assertIn("Knowledge Entry 1:", context) |
| self.assertIn("Q: Test question?", context) |
| self.assertIn("A: Test answer.", context) |
| self.assertIn("-" * 40, context) |
|
|
| def test_process_context_missing_metadata(self): |
| """Test context processing with missing metadata fields""" |
| mock_result = Mock() |
| mock_result.metadata = {} |
|
|
| context, source_ids, knowledge_pairs = _process_context_impl( |
| [mock_result], [0.9], max_results=1 |
| ) |
|
|
| |
| self.assertIn("N/A", context) |
| self.assertEqual(source_ids[0], "N/A") |
|
|
| @patch("src.vector_store.get_knowledge_base_data") |
| @patch("src.vector_store.chromadb.PersistentClient") |
| @patch("src.vector_store.Chroma") |
| def test_initialize_vector_store_new_collection( |
| self, mock_chroma_class, mock_client_class, mock_get_kb |
| ): |
| """Test initializing vector store with new collection""" |
| |
| mock_get_kb.return_value = ( |
| ["doc1", "doc2"], |
| [{"id": "1"}, {"id": "2"}], |
| ["id1", "id2"], |
| ) |
|
|
| |
| mock_client = Mock() |
| mock_client_class.return_value = mock_client |
| |
| |
| mock_client.get_collection.side_effect = Exception("Collection not found") |
| |
| |
| mock_collection = Mock() |
| mock_client.create_collection.return_value = mock_collection |
|
|
| |
| mock_vector_store = Mock() |
| mock_retriever = Mock() |
| mock_vector_store.as_retriever.return_value = mock_retriever |
| mock_chroma_class.return_value = mock_vector_store |
|
|
| |
| from src.vector_store import initialize_vector_store |
|
|
| collection, vector_store, retriever = initialize_vector_store() |
|
|
| |
| mock_client.create_collection.assert_called_once() |
| mock_collection.add.assert_called_once() |
|
|
| |
| self.assertEqual(vector_store, mock_vector_store) |
| self.assertEqual(retriever, mock_retriever) |
|
|
| @patch("src.vector_store.get_knowledge_base_data") |
| @patch("src.vector_store.chromadb.PersistentClient") |
| @patch("src.vector_store.Chroma") |
| def test_initialize_vector_store_existing_collection( |
| self, mock_chroma_class, mock_client_class, mock_get_kb |
| ): |
| """Test initializing vector store with existing collection""" |
| |
| mock_get_kb.return_value = ( |
| ["doc1", "doc2"], |
| [{"id": "1"}, {"id": "2"}], |
| ["id1", "id2"], |
| ) |
|
|
| |
| mock_client = Mock() |
| mock_client_class.return_value = mock_client |
| |
| |
| mock_collection = Mock() |
| mock_client.get_collection.return_value = mock_collection |
|
|
| |
| mock_vector_store = Mock() |
| mock_retriever = Mock() |
| mock_vector_store.as_retriever.return_value = mock_retriever |
| mock_chroma_class.return_value = mock_vector_store |
|
|
| |
| from src.vector_store import initialize_vector_store |
|
|
| collection, vector_store, retriever = initialize_vector_store() |
|
|
| |
| mock_client.get_collection.assert_called_once() |
| mock_client.create_collection.assert_not_called() |
|
|
| |
| self.assertEqual(collection, mock_collection) |
| self.assertEqual(vector_store, mock_vector_store) |
| self.assertEqual(retriever, mock_retriever) |
|
|
| @patch("src.vector_store.get_knowledge_base_data") |
| @patch("src.vector_store.chromadb.PersistentClient") |
| def test_initialize_vector_store_failure(self, mock_client_class, mock_get_kb): |
| """Test initialize_vector_store handles errors properly""" |
| |
| mock_get_kb.return_value = (["doc1"], [{"id": "1"}], ["id1"]) |
|
|
| |
| mock_client_class.side_effect = Exception("Database connection failed") |
|
|
| |
| from src.vector_store import initialize_vector_store |
|
|
| with self.assertRaises(Exception) as context: |
| initialize_vector_store() |
|
|
| self.assertIn("Database connection failed", str(context.exception)) |
|
|
|
|
| if __name__ == "__main__": |
| unittest.main() |
|
|