Kelvin-programmer/rag-chatbot
0
1"""Tests for the VectorStore class."""2 3import os4 5import pytest6 7from src.vector_store import VectorStore8 9EMBEDDING_MODEL = "all-MiniLM-L12-v2"10 11 12@pytest.fixture13def store(temp_dir):14 return VectorStore(model_name=EMBEDDING_MODEL, persist_dir=temp_dir)15 16 17class TestVectorStoreInit:18 def test_dimension_positive(self, store):19 assert store.dimension > 020 21 def test_starts_empty(self, store):22 assert store.count == 023 assert store.documents == []24 25 26class TestAddDocuments:27 def test_add_returns_count(self, store, sample_texts):28 assert store.add_documents(sample_texts) == len(sample_texts)29 30 def test_count_after_add(self, store, sample_texts):31 store.add_documents(sample_texts)32 assert store.count == len(sample_texts)33 34 def test_add_empty_list(self, store):35 assert store.add_documents([]) == 036 assert store.count == 037 38 def test_metadata_stored(self, store):39 docs = ["hello world"]40 meta = [{"source": "test.pdf", "page": 1}]41 store.add_documents(docs, metadata=meta)42 results = store.search("hello", top_k=1)43 assert results[0]["metadata"] == meta[0]44 45 46class TestSearch:47 def test_returns_relevant_results(self, store, sample_texts):48 store.add_documents(sample_texts)49 results = store.search("What is machine learning?", top_k=2)50 assert len(results) == 251 assert any("machine learning" in r["text"].lower() for r in results)52 53 def test_empty_store(self, store):54 assert store.search("anything") == []55 56 def test_respects_top_k(self, store, sample_texts):57 store.add_documents(sample_texts)58 assert len(store.search("programming", top_k=1)) == 159 60 def test_result_has_score(self, store, sample_texts):61 store.add_documents(sample_texts)62 results = store.search("Python", top_k=1)63 assert "score" in results[0]64 assert 0 < results[0]["score"] <= 165 66 67class TestPersistence:68 def test_save_and_load(self, store, sample_texts, temp_dir):69 store.add_documents(sample_texts)70 store.save()71 72 new_store = VectorStore(model_name=EMBEDDING_MODEL, persist_dir=temp_dir)73 assert new_store.load() is True74 assert new_store.count == len(sample_texts)75 assert new_store.documents == store.documents76 77 def test_load_nonexistent(self, temp_dir):78 store = VectorStore(79 model_name=EMBEDDING_MODEL,80 persist_dir=os.path.join(temp_dir, "nonexistent"),81 )82 assert store.load() is False83 84 85class TestClear:86 def test_clear_resets(self, store, sample_texts):87 store.add_documents(sample_texts)88 store.clear()89 assert store.count == 090 assert store.documents == []91 