Team Ai
Apppublic

Underground-Digital/Workflow-Engine

sourceHugging Faceupdated 2y agoView on Hugging Face
0likes
test_vector_store.py96 linesDownload Raw Back to vdb
1import uuid2from unittest.mock import MagicMock3 4import pytest5 6from core.rag.models.document import Document7from extensions import ext_redis8from models.dataset import Dataset9 10 11def get_example_text() -> str:12    return "test_text"13 14 15def get_example_document(doc_id: str) -> Document:16    doc = Document(17        page_content=get_example_text(),18        metadata={19            "doc_id": doc_id,20            "doc_hash": doc_id,21            "document_id": doc_id,22            "dataset_id": doc_id,23        },24    )25    return doc26 27 28@pytest.fixture29def setup_mock_redis() -> None:30    # get31    ext_redis.redis_client.get = MagicMock(return_value=None)32 33    # set34    ext_redis.redis_client.set = MagicMock(return_value=None)35 36    # lock37    mock_redis_lock = MagicMock()38    mock_redis_lock.__enter__ = MagicMock()39    mock_redis_lock.__exit__ = MagicMock()40    ext_redis.redis_client.lock = mock_redis_lock41 42 43class AbstractVectorTest:44    def __init__(self):45        self.vector = None46        self.dataset_id = str(uuid.uuid4())47        self.collection_name = Dataset.gen_collection_name_by_id(self.dataset_id) + "_test"48        self.example_doc_id = str(uuid.uuid4())49        self.example_embedding = [1.001 * i for i in range(128)]50 51    def create_vector(self) -> None:52        self.vector.create(53            texts=[get_example_document(doc_id=self.example_doc_id)],54            embeddings=[self.example_embedding],55        )56 57    def search_by_vector(self):58        hits_by_vector: list[Document] = self.vector.search_by_vector(query_vector=self.example_embedding)59        assert len(hits_by_vector) == 160        assert hits_by_vector[0].metadata["doc_id"] == self.example_doc_id61 62    def search_by_full_text(self):63        hits_by_full_text: list[Document] = self.vector.search_by_full_text(query=get_example_text())64        assert len(hits_by_full_text) == 165        assert hits_by_full_text[0].metadata["doc_id"] == self.example_doc_id66 67    def delete_vector(self):68        self.vector.delete()69 70    def delete_by_ids(self, ids: list[str]):71        self.vector.delete_by_ids(ids=ids)72 73    def add_texts(self) -> list[str]:74        batch_size = 10075        documents = [get_example_document(doc_id=str(uuid.uuid4())) for _ in range(batch_size)]76        embeddings = [self.example_embedding] * batch_size77        self.vector.add_texts(documents=documents, embeddings=embeddings)78        return [doc.metadata["doc_id"] for doc in documents]79 80    def text_exists(self):81        assert self.vector.text_exists(self.example_doc_id)82 83    def get_ids_by_metadata_field(self):84        with pytest.raises(NotImplementedError):85            self.vector.get_ids_by_metadata_field(key="key", value="value")86 87    def run_all_tests(self):88        self.create_vector()89        self.search_by_vector()90        self.search_by_full_text()91        self.text_exists()92        self.get_ids_by_metadata_field()93        added_doc_ids = self.add_texts()94        self.delete_by_ids(added_doc_ids)95        self.delete_vector()96