diff --git a/datastew/__init__.py b/datastew/__init__.py index e5a03de..e94a387 100644 --- a/datastew/__init__.py +++ b/datastew/__init__.py @@ -7,11 +7,9 @@ Concept, Mapping, Terminology, - base, model, pagination, postgresql, - sqllite, ) from .visualisation import ( bar_chart_average_acc_two_distributions, @@ -26,11 +24,9 @@ "mapping", "ols", "parsing", - "base", "model", "pagination", "postgresql", - "sqllite", "EmbeddingModel", "GPT4Adapter", "HuggingFaceAdapter", diff --git a/datastew/process/jsonl_adapter.py b/datastew/process/jsonl_adapter.py index 3aa4c75..2e2a836 100644 --- a/datastew/process/jsonl_adapter.py +++ b/datastew/process/jsonl_adapter.py @@ -7,7 +7,7 @@ from tqdm import tqdm from datastew.embedding import Vectorizer -from datastew.repository import PostgreSQLRepository, SQLLiteRepository +from datastew.repository import PostgreSQLRepository from datastew.repository.model import Concept, Mapping, Terminology @@ -69,7 +69,7 @@ def _flush_to_file(self, file_path: str): self._buffer.clear() - def from_repository(self, repository: Union[PostgreSQLRepository, SQLLiteRepository]): + def from_repository(self, repository: PostgreSQLRepository): """Export all records from a PostgreSQLRepository to JSONL files :param repository: Active database repository instace. diff --git a/datastew/process/ols.py b/datastew/process/ols.py index ba625b7..44d83eb 100644 --- a/datastew/process/ols.py +++ b/datastew/process/ols.py @@ -5,8 +5,8 @@ import requests from datastew.embedding import Vectorizer -from datastew.repository.base import BaseRepository from datastew.repository.model import Concept, Mapping, Terminology +from datastew.repository.postgresql import PostgreSQLRepository class OLSTerminologyImportTask: @@ -83,7 +83,7 @@ def __process_page(self, page: int) -> Tuple[List[Concept], List[Mapping]]: logging.error(f"Failed to fetch concepts and descriptions from OLS for page {page}: {str(e)}") return [], [] - def process_to_repository(self, repository: BaseRepository): + def process_to_repository(self, repository: PostgreSQLRepository): """ Fetches concepts and descriptions from the OLS API and stores them in a repository. diff --git a/datastew/repository/__init__.py b/datastew/repository/__init__.py index 3df8a0b..78a3063 100644 --- a/datastew/repository/__init__.py +++ b/datastew/repository/__init__.py @@ -1,5 +1,4 @@ from .model import Concept, Mapping, Terminology from .postgresql import PostgreSQLRepository -from .sqllite import SQLLiteRepository -__all__ = ["Terminology", "Concept", "Mapping", "SQLLiteRepository", "PostgreSQLRepository"] +__all__ = ["Terminology", "Concept", "Mapping", "PostgreSQLRepository"] diff --git a/datastew/repository/base.py b/datastew/repository/base.py deleted file mode 100644 index d8138dd..0000000 --- a/datastew/repository/base.py +++ /dev/null @@ -1,124 +0,0 @@ -import logging -from abc import ABC, abstractmethod -from typing import List, Optional, Sequence, Union - -from datastew.embedding import Vectorizer -from datastew.process.parsing import DataDictionarySource -from datastew.repository.model import Concept, Mapping, MappingResult, Terminology -from datastew.repository.pagination import Page - -logger = logging.getLogger(__name__) - - -class BaseRepository(ABC): - def __init__(self, vectorizer: Vectorizer = Vectorizer()): - self.vectorizer = vectorizer - - @abstractmethod - def store(self, model_object_instance: Union[Terminology, Concept, Mapping]): - """Store a single model object instance.""" - pass - - @abstractmethod - def store_all(self, model_object_instances: List[Union[Terminology, Concept, Mapping]]): - """Store multiple model object instances.""" - pass - - @abstractmethod - def get_concept(self, concept_id: str) -> Concept: - """Retrieve a Concept by ID from the database.""" - pass - - @abstractmethod - def get_concepts(self, terminology_name: Optional[str] = None, offset: int = 0, limit: int = 100) -> Page[Concept]: - """Retrieve all concepts from the database.""" - pass - - @abstractmethod - def get_terminology(self, terminology_name: str) -> Terminology: - """Retrieve a Terminology by name from the database.""" - pass - - @abstractmethod - def get_all_terminologies(self) -> List[Terminology]: - """Retrieve all terminologies from the database.""" - pass - - @abstractmethod - def get_mappings( - self, - terminology_name: Optional[str] = None, - sentence_embedder: Optional[str] = None, - limit: int = 1000, - offset: int = 0, - ) -> Page[Mapping]: - """Get all embeddings up to a limit""" - pass - - @abstractmethod - def get_all_sentence_embedders(self) -> List[str]: - pass - - @abstractmethod - def get_closest_mappings( - self, - embedding: Sequence[float], - similarities: bool = False, - terminology_name: Optional[str] = None, - sentence_embedder: Optional[str] = None, - limit=5, - ) -> Union[List[Mapping], List[MappingResult]]: - """Get the closest mappings based on embedding.""" - pass - - @abstractmethod - def shut_down(self): - """Shut down the repository.""" - pass - - @abstractmethod - def clear_all(self): - """Clear all entries in the database.""" - pass - - def import_data_dictionary(self, data_dictionary: DataDictionarySource, terminology_name: str): - """Imports a data dictionary, generating concepts and embeddings, and stores them in the database. - - :param data_dictionary: Source of variable descriptions and metadata. - :param terminology_name: Name of the terminology being imported. - :raises RuntimeError: If the import or transformation fails. - """ - try: - objects = self._parse_data_dictionary(data_dictionary, terminology_name) - self.store_all(objects) - except Exception as e: - logger.exception("Failed to import data dictionary.") - raise RuntimeError(f"Failed to import data dictionary source: {e}") - - @abstractmethod - def import_from_jsonl(self, jsonl_path: str, object_type: str, chunk_size: int = 100): - pass - - def _parse_data_dictionary( - self, data_dictionary: DataDictionarySource, terminology_name: str - ) -> List[Union[Concept, Mapping, Terminology]]: - df = data_dictionary.to_dataframe() - descriptions = df["description"].tolist() - vectorizer_name = self.vectorizer.model_name - variable_to_embedding = data_dictionary.get_embeddings(self.vectorizer) - - terminology = Terminology(name=terminology_name, id=terminology_name) - objects: List[Union[Concept, Mapping, Terminology]] = [terminology] - - for variable, description in zip(variable_to_embedding.keys(), descriptions): - concept_id = f"{terminology_name}:{variable}" - concept = Concept(terminology=terminology, pref_label=variable, concept_identifier=concept_id) - mapping = Mapping( - concept=concept, - text=description, - embedding=variable_to_embedding[variable], - sentence_embedder=vectorizer_name, - ) - objects.extend([concept, mapping]) - - return objects diff --git a/datastew/repository/postgresql.py b/datastew/repository/postgresql.py index 0ee0cc5..53b57ea 100644 --- a/datastew/repository/postgresql.py +++ b/datastew/repository/postgresql.py @@ -7,14 +7,14 @@ from datastew.embedding import Vectorizer from datastew.exceptions import ObjectStorageError -from datastew.repository.base import BaseRepository +from datastew.process.parsing import DataDictionarySource from datastew.repository.model import Base, Concept, Mapping, MappingResult, Terminology from datastew.repository.pagination import Page logger = logging.getLogger(__name__) -class PostgreSQLRepository(BaseRepository): +class PostgreSQLRepository: def __init__( self, connection_string: str, @@ -33,7 +33,7 @@ def __init__( :param pool_timeout: The maximum time (in seconds) to wait for a connection from the pool before raising an exception. """ - super().__init__(vectorizer) + self.vectorizer = vectorizer self.engine = create_engine( connection_string, pool_size=pool_size, max_overflow=max_overflow, pool_timeout=pool_timeout ) @@ -197,6 +197,20 @@ def clear_all(self): self.session.query(Terminology).delete() self.session.commit() + def import_data_dictionary(self, data_dictionary: DataDictionarySource, terminology_name: str): + """Imports a data dictionary, generating concepts and embeddings, and stores them in the database. + + :param data_dictionary: Source of variable descriptions and metadata. + :param terminology_name: Name of the terminology being imported. + :raises RuntimeError: If the import or transformation fails. + """ + try: + objects = self._parse_data_dictionary(data_dictionary, terminology_name) + self.store_all(objects) + except Exception as e: + logger.exception("Failed to import data dictionary.") + raise RuntimeError(f"Failed to import data dictionary source: {e}") + def import_from_jsonl( self, jsonl_path: str, object_type: Literal["terminology", "concept", "mapping"], chunk_size: int = 100 ): @@ -305,3 +319,27 @@ def _validate_required_fields( def _initialize_pgvector(self): with self.engine.begin() as conn: conn.execute(text("CREATE EXTENSION IF NOT EXISTS vector")) + + def _parse_data_dictionary( + self, data_dictionary: DataDictionarySource, terminology_name: str + ) -> List[Union[Concept, Mapping, Terminology]]: + df = data_dictionary.to_dataframe() + descriptions = df["description"].tolist() + vectorizer_name = self.vectorizer.model_name + variable_to_embedding = data_dictionary.get_embeddings(self.vectorizer) + + terminology = Terminology(name=terminology_name, id=terminology_name) + objects: List[Union[Concept, Mapping, Terminology]] = [terminology] + + for variable, description in zip(variable_to_embedding.keys(), descriptions): + concept_id = f"{terminology_name}:{variable}" + concept = Concept(terminology=terminology, pref_label=variable, concept_identifier=concept_id) + mapping = Mapping( + concept=concept, + text=description, + embedding=variable_to_embedding[variable], + sentence_embedder=vectorizer_name, + ) + objects.extend([concept, mapping]) + + return objects diff --git a/datastew/repository/sqllite.py b/datastew/repository/sqllite.py deleted file mode 100644 index 28c8bd8..0000000 --- a/datastew/repository/sqllite.py +++ /dev/null @@ -1,320 +0,0 @@ -import json -import logging -from typing import Any, Dict, List, Literal, Optional, Union - -import numpy as np -from sqlalchemy import create_engine, func -from sqlalchemy.orm import joinedload, sessionmaker -from sqlalchemy.pool import StaticPool - -from datastew.embedding import Vectorizer -from datastew.exceptions import ObjectStorageError -from datastew.repository.base import BaseRepository -from datastew.repository.model import Base, Concept, Mapping, MappingResult, Terminology -from datastew.repository.pagination import Page - -logger = logging.getLogger(__name__) - - -class SQLLiteRepository(BaseRepository): - - def __init__( - self, - mode: str = "memory", - path: Optional[str] = None, - vectorizer: Vectorizer = Vectorizer(), - ): - """Initializes the repository with a SQLite backend. - - :param mode: Storage mode control, defaults to "memory". - :param path: File path to SQLite DB when mode is "disk", defaults to None. - :param vectorizer: An instance of Vectorizer for generating embeddings, defaults to Vectorizer(). - :raises ValueError: Undefined DB mode. - """ - super().__init__(vectorizer) - if mode == "disk": - self.engine = create_engine(f"sqlite:///{path}") - # for tests - elif mode == "memory": - self.engine = create_engine( - "sqlite:///:memory:", connect_args={"check_same_thread": False}, poolclass=StaticPool - ) - else: - raise ValueError(f"DB mode {mode} is not defined. Use either disk or memory.") - Base.metadata.create_all(self.engine) - Session = sessionmaker(bind=self.engine, autoflush=False) - self.session = Session() - - def store(self, model_object_instance: Union[Terminology, Concept, Mapping]): - """Stores a single Terminology, Concept, or Mapping object in the database. - - :param model_object_instance: An instance of Terminology, Concept, or Mapping. - :raises ObjectStorageError: If the object cannot be stored (e.g., due to DB errors). - """ - try: - self.session.merge(model_object_instance) - self.session.commit() - except Exception as e: - self.session.rollback() - logger.exception("Failed to store object.") - raise ObjectStorageError("Failed to store object in the database.", e) - - def store_all(self, model_object_instances: List[Union[Terminology, Concept, Mapping]]): - """Stores a list of Terminology, Concept, or Mapping objects in the database. - - :param model_object_instances: List of model objects to store. - """ - for obj in model_object_instances: - self.store(obj) - - def get_concept(self, concept_id: str) -> Concept: - """Retrieves a Concept by its ID. - - :param concept_id: ID of the Concept. - :raises ValueError: If no Concept with given ID is found. - :return: Concept object. - """ - concept = self.session.query(Concept).filter_by(concept_identifier=concept_id).first() - if concept is None: - raise ValueError(f"No Concept found with ID: {concept_id}") - return concept - - def get_concepts(self, terminology_name: Optional[str] = None, offset: int = 0, limit: int = 100) -> Page[Concept]: - """Retrieves all concepts from the database. - - :return: All stored Concept objects. - """ - query = self.session.query(Concept).options(joinedload(Concept.terminology)) - - if terminology_name: - query = query.join(Concept.terminology).filter(Terminology.name == terminology_name) - - total_count = query.with_entities(func.count()).scalar() - concepts = query.offset(offset).limit(limit).all() - return Page[Concept](items=concepts, limit=limit, offset=offset, total_count=total_count) - - def get_terminology(self, terminology_name: str) -> Terminology: - """Retrieves a Terminology objects by its name. - - :param terminology_name: Name of the terminology. - :raises ValueError: If no terminology with the given name is found. - :return: Terminology object. - """ - terminology = self.session.query(Terminology).filter_by(name=terminology_name).first() - if terminology is None: - raise ValueError(f"No Terminology found with name: {terminology_name}") - return terminology - - def get_all_terminologies(self) -> List[Terminology]: - """Retrieves all terminologies from the database. - - :return: All stored Terminology objects. - """ - return self.session.query(Terminology).all() - - def get_mappings( - self, - terminology_name: Optional[str] = None, - sentence_embedder: Optional[str] = None, - limit: int = 1000, - offset: int = 0, - ) -> Page[Mapping]: - """Retrieves a paginated list of mappings, optionally filtered by terminology name and/or sentence embedder. - - :param terminology_name: Name of the terminology to filter by, defaults to None - :param sentence_embedder: Name of the sentence embedding model to filter by, defaults to None - :param limit: Maximum number of results to return, defaults to 1000 - :param offset: Number of items to skip, defaults to 0 - :return: A paginated result containing mappings and metadata. - """ - query = self.session.query(Mapping) - - if terminology_name: - query = query.join(Concept).join(Terminology).filter(Terminology.name == terminology_name) - - if sentence_embedder: - query = query.filter(Mapping.sentence_embedder == sentence_embedder) - - total_count = query.count() - - if total_count == 0: - return Page(items=[], limit=limit, offset=offset, total_count=total_count) - - items = query.offset(offset).limit(limit).all() - - return Page(items=items, limit=limit, offset=offset, total_count=total_count) - - def get_all_sentence_embedders(self) -> List[str]: - """Retrieves all distinct sentence embedder names used in the mappings. - - :return: Unique sentence embedder identifiers. - """ - return [embedder for embedder, in self.session.query(Mapping.sentence_embedder).distinct().all()] - - def get_closest_mappings( - self, - embedding: List[float], - similarities: bool = True, - terminology_name: Optional[str] = None, - sentence_embedder: Optional[str] = None, - limit: int = 5, - ) -> Union[List[Mapping], List[MappingResult]]: - """Finds the closest mappings by cosine similarity to a given embedding, optionally filtered. - - :param embedding: The target embedding vector to compare against. - :param similarities: If True, returns MappingResult objects with similarity scores, defaults to True. - :param terminology_name: Filter by terminology name, defaults to None. - :param sentence_embedder: Filter by sentence embedder name, defaults to None. - :param limit: Maximum number of results to return, defaults to 5. - :return: Closest mappings, with or without similarity scores. - """ - query = self.session.query(Mapping) - - if terminology_name: - query = query.join(Concept).join(Terminology).filter(Terminology.name == terminology_name) - - if sentence_embedder: - query = query.filter(Mapping.sentence_embedder == sentence_embedder) - - mappings = query.all() - - if not mappings: - return [] - - all_embeddings = np.array([mapping.embedding for mapping in mappings]) - target_embedding = np.array(embedding) - - if similarities: - denominator = np.linalg.norm(all_embeddings, axis=1) * np.linalg.norm(target_embedding) - # Substitute denominator with 1e-10 in case the one of the either norms is 0 - denominator = np.where(denominator == 0, 1e-10, denominator) - similarity = np.dot(all_embeddings, target_embedding) / denominator - sorted_indices = np.argsort(similarity)[::-1] - results = [MappingResult(mapping=mappings[i], similarity=similarity[i]) for i in sorted_indices[:limit]] - return results - - return mappings[:limit] - - def shut_down(self): - """ - Closes the SQLAlchemy session and releases database resources. - """ - self.session.close() - - def clear_all(self): - """Deletes all Terminology, Concept, and Mapping entries from the database. - - This method is primarily intended for test environments to ensure a clean - state before or after test execution. It performs bulk deletions in the - correct dependency order (Mappings → Concepts → Terminologies) and commits - the changes. - """ - self.session.query(Mapping).delete() - self.session.query(Concept).delete() - self.session.query(Terminology).delete() - self.session.commit() - - def import_from_jsonl( - self, jsonl_path: str, object_type: Literal["terminology", "concept", "mapping"], chunk_size: int = 100 - ): - """Imports data from a JSONL file and stores it in the database in chunks. - - :param jsonl_path: Path to the JSONL file containing the data to be imported. - :param object_type: Literal specifying the object type, must be "terminology", "concept", or "mapping". - :param chunk_size: Number of objects to store in a single batch, defaults to 100. - :raises ValueError: If the JSON is malformed or required fields are missing. - :raises ValueError: If the provided `object_type` is unsupported. - :raises RuntimeError: If the file cannot be found. - :raises RuntimeError: If a general I/O or database error occurs during import. - """ - buffer = [] - - try: - with open(jsonl_path, "r", encoding="utf-8") as file: - for idx, line in enumerate(file): - line = line.strip() - if not line: - continue - try: - data = json.loads(line) - except json.JSONDecodeError as e: - raise ValueError(f"Invalid JSON on line {idx + 1}: {e}") - - obj = self._deserialize_object(object_type, data) - buffer.append(obj) - - if len(buffer) >= chunk_size: - self.store_all(buffer) - buffer = [] - - if buffer: - self.store_all(buffer) - except ValueError: - raise - except FileNotFoundError: - raise RuntimeError(f"File not found: {jsonl_path}") - except Exception as e: - raise RuntimeError(f"Error while importing {object_type}: {e}") - - def _deserialize_object( - self, object_type: Literal["terminology", "concept", "mapping"], data: Dict[str, Any] - ) -> Union[Terminology, Concept, Mapping]: - """Deserializes a JSON object into an SQLAlchemy model instance, resolving any required relationships. - - :param object_type: The type of object to deserialize. - :param data: The dictionary representing the object, as loaded from a JSONL line. - :raises ValueError: If a related object (e.g., a referenced concept or terminology) cannot be found. - :raises ValueError: If required attributes are missing. - :raises ValueError: If the object_type is not one of the supported values. - :return: An instance of the appropriate SQLAlchemy model. - """ - if object_type == "terminology": - # Validate required keys - self._validate_required_fields(data, ["id", "name"], object_type) - return Terminology(**data) - - elif object_type == "concept": - # Validate required keys - self._validate_required_fields(data, ["terminology_id", "pref_label", "concept_identifier"], object_type) - - terminology = self.session.get(Terminology, data["terminology_id"]) - if not terminology: - raise ValueError(f"Terminology with ID {data['terminology_id']} not found") - - return Concept( - terminology=terminology, - pref_label=data["pref_label"], - concept_identifier=data["concept_identifier"], - ) - - elif object_type == "mapping": - # Validate required keys - self._validate_required_fields(data, ["concept_identifier", "text"], object_type) - - concept = self.session.get(Concept, data["concept_identifier"]) - if not concept: - raise ValueError(f"Concept with ID {data['concept_identifier']} not found") - - embedding = data.get("embedding") - sentence_embedder = data.get("sentence_embedder") - - if (embedding is None or sentence_embedder is None) and self.vectorizer: - embedding = self.vectorizer.get_embedding(data["text"]) - sentence_embedder = self.vectorizer.model_name - - return Mapping( - concept=concept, - text=data["text"], - embedding=embedding, - sentence_embedder=sentence_embedder, - ) - - else: - raise ValueError(f"Unsupported object_type: {object_type}") - - def _validate_required_fields( - self, data: Dict[str, Any], required_keys: List[str], object_type: Literal["terminology", "concept", "mapping"] - ): - for key in required_keys: - if key not in data: - raise ValueError(f"Missing required field '{key}' for {object_type}") diff --git a/datastew/visualisation.py b/datastew/visualisation.py index 19c3277..0a902ec 100644 --- a/datastew/visualisation.py +++ b/datastew/visualisation.py @@ -10,7 +10,7 @@ from datastew.embedding import Vectorizer from datastew.process.parsing import DataDictionarySource -from datastew.repository.base import BaseRepository +from datastew.repository.postgresql import PostgreSQLRepository def enrichment_plot(acc_gpt, acc_mpnet, acc_fuzzy, title, save_plot=False, save_dir="resources/results/plots"): @@ -88,7 +88,7 @@ def bar_chart_average_acc_two_distributions( def get_plot_for_current_database_state( - repository: BaseRepository, + repository: PostgreSQLRepository, terminology: Optional[str] = None, sentence_embedder: Optional[str] = None, limit: int = 1000, diff --git a/tests/test_postgreqsl_repository.py b/tests/test_postgreqsl_repository.py deleted file mode 100644 index aee0633..0000000 --- a/tests/test_postgreqsl_repository.py +++ /dev/null @@ -1,17 +0,0 @@ -import os - -from datastew.process.jsonl_adapter import SQLJsonlConverter -from datastew.repository.postgresql import PostgreSQLRepository -from tests.base_repository_test_setup import BaseRepositoryTestSetup - - -class TestPostgreSQLRepository(BaseRepositoryTestSetup): - __test__ = True - POSTGRES_TEST_URL = os.getenv("TEST_POSTGRES_URI", "postgresql://testuser:testpass@localhost/testdb") - - @classmethod - def setUpClass(cls): - super().setUpClass() - cls.repo_args = (cls.POSTGRES_TEST_URL, cls.vectorizer1) - cls.repository = PostgreSQLRepository(*cls.repo_args) - cls.jsonl_converter = SQLJsonlConverter(dest_dir="test_export") diff --git a/tests/base_repository_test_setup.py b/tests/test_postgresql_repository.py similarity index 94% rename from tests/base_repository_test_setup.py rename to tests/test_postgresql_repository.py index e4bd0e9..b23ce2d 100644 --- a/tests/base_repository_test_setup.py +++ b/tests/test_postgresql_repository.py @@ -6,15 +6,13 @@ from datastew.process.jsonl_adapter import SQLJsonlConverter from datastew.process.parsing import DataDictionarySource from datastew.repository import Concept, Mapping, Terminology -from datastew.repository.base import BaseRepository +from datastew.repository.postgresql import PostgreSQLRepository -class BaseRepositoryTestSetup(unittest.TestCase): - """Base class for setting up test data and shared tests for all repository backends.""" +class TestPostgreSQLRepository(unittest.TestCase): + """Tests for the PostgreSQL repository backend.""" - __test__ = False - repository: BaseRepository - jsonl_converter: SQLJsonlConverter + POSTGRES_TEST_URL = os.getenv("TEST_POSTGRES_URI", "postgresql://testuser:testpass@localhost/testdb") TEST_CONCEPTS = [ ("Diabetes mellitus (disorder)", "Concept ID: 11893007", "v1"), @@ -32,7 +30,7 @@ class BaseRepositoryTestSetup(unittest.TestCase): @classmethod def setUpClass(cls): - """Shared setup for vectorizers and paths.""" + """Setup for vectorizers, paths, and PostgreSQL repository.""" cls.TEST_DIR_PATH = os.path.dirname(os.path.realpath(__file__)) cls.vectorizer1 = Vectorizer("sentence-transformers/all-mpnet-base-v2") cls.vectorizer2 = Vectorizer("FremyCompany/BioLORD-2023") @@ -40,6 +38,10 @@ def setUpClass(cls): cls.model_name2 = cls.vectorizer2.model_name cls.test_text = "The flu" + cls.repo_args = (cls.POSTGRES_TEST_URL, cls.vectorizer1) + cls.repository = PostgreSQLRepository(*cls.repo_args) + cls.jsonl_converter = SQLJsonlConverter(dest_dir="test_export") + @classmethod def tearDownClass(cls): """Ensure repository is properly closed""" diff --git a/tests/test_sqllite_import.py b/tests/test_sqllite_import.py deleted file mode 100644 index e131309..0000000 --- a/tests/test_sqllite_import.py +++ /dev/null @@ -1,124 +0,0 @@ -import json -import os -import shutil -import tempfile -from typing import Any, Dict, List, Literal -from unittest import TestCase - -from datastew.repository.sqllite import SQLLiteRepository - - -class TestWeaviateRepositoryImport(TestCase): - - def setUp(self) -> None: - self.repository = SQLLiteRepository("disk", "sqlite_db") - self.repository.clear_all() - self.temp_dir = tempfile.mkdtemp() - - # Sample data for JSONL files - self.data_files = { - "terminology": [{"id": "import_test", "name": "import_test"}], - "concept": [ - {"concept_identifier": "import_test:G", "pref_label": "G", "terminology_id": "import_test"}, - {"concept_identifier": "import_test:H", "pref_label": "H", "terminology_id": "import_test"}, - ], - "mapping": [ - { - "text": "pancreas", - "concept_identifier": "import_test:G", - "embedding": [0.048744574189186096, -0.0035385489463806152], - "sentence_embedder": "sentence-transformers/all-mpnet-base-v2", - }, - { - "text": "liver", - "concept_identifier": "import_test:H", - "embedding": [0.1, -0.2], - "sentence_embedder": "sentence-transformers/all-mpnet-base-v2", - }, - ], - } - - # Write data to JSONL files - for key, data in self.data_files.items(): - self.write_jsonl(os.path.join(self.temp_dir, f"{key}.jsonl"), data) - - @staticmethod - def write_jsonl(file_path: str, data: List[Dict[str, Any]]): - """Write data to a JSONL file.""" - with open(file_path, "w") as file: - for obj in data: - json.dump(obj, file) - file.write("\n") - - def import_data(self, data_types: List[Literal["terminology", "concept", "mapping"]]): - """Helper method to import multiple data types.""" - for data_type in data_types: - file_path = os.path.join(self.temp_dir, f"{data_type}.jsonl") - self.repository.import_from_jsonl(file_path, data_type) - - def tearDown(self) -> None: - shutil.rmtree(self.temp_dir, ignore_errors=True) - self.repository.shut_down() - - def test_import_terminology(self): - self.import_data(["terminology"]) - terminology = self.repository.get_all_terminologies() - - self.assertEqual(len(terminology), 1) - with self.subTest("Terminology ID"): - self.assertEqual(terminology[0].id, "import_test") - with self.subTest("Terminology Name"): - self.assertEqual(terminology[0].name, "import_test") - - def test_import_concepts(self): - self.import_data(["terminology", "concept"]) - concepts = self.repository.get_concepts(limit=5, offset=0).items - - self.assertEqual(len(concepts), 2) - - for concept in concepts: - with self.subTest(f"Concept Properties: {concept.pref_label}"): - self.assertIn(concept.pref_label, ["G", "H"]) - self.assertIn(concept.concept_identifier, ["import_test:G", "import_test:H"]) - with self.subTest(f"Terminology Reference for {concept.pref_label}"): - self.assertEqual(concept.terminology.name, "import_test") - - def test_import_mappings(self): - self.import_data(["terminology", "concept", "mapping"]) - mappings = self.repository.get_mappings(limit=10, offset=0).items - - self.assertEqual(len(mappings), 2) - - for mapping in mappings: - with self.subTest(f"Mapping Text for {mapping.text}"): - self.assertIn(mapping.text, ["pancreas", "liver"]) - with self.subTest(f"Sentence Embedder for {mapping.text}"): - self.assertEqual(mapping.sentence_embedder, "sentence-transformers/all-mpnet-base-v2") - with self.subTest(f"Vector Length for {mapping.text}"): - self.assertEqual(len(mapping.embedding), 2) - with self.subTest(f"Concept Reference for Mapping {mapping.text}"): - expected_label = "G" if mapping.text == "pancreas" else "H" - self.assertEqual(mapping.concept.pref_label, expected_label) - - def test_import_invalid_jsonl(self): - invalid_file = os.path.join(self.temp_dir, "invalid.jsonl") - with open(invalid_file, "w") as file: - file.write("{ invalid jsonl }") - - with self.assertRaises(ValueError): - self.repository.import_from_jsonl(invalid_file, "terminology") - - def test_import_missing_id(self): - file_path = os.path.join(self.temp_dir, "missing_id.jsonl") - self.write_jsonl(file_path, [{"name": "missing_id"}]) - - with self.assertRaises(ValueError): - self.repository.import_from_jsonl(file_path, "terminology") - - def test_import_empty_file(self): - empty_file = os.path.join(self.temp_dir, "empty.jsonl") - open(empty_file, "w").close() - - self.repository.import_from_jsonl(empty_file, "terminology") - terminology = self.repository.get_all_terminologies() - self.assertEqual(len(terminology), 0) diff --git a/tests/test_sqllite_repository.py b/tests/test_sqllite_repository.py deleted file mode 100644 index 0ef94f4..0000000 --- a/tests/test_sqllite_repository.py +++ /dev/null @@ -1,14 +0,0 @@ -from datastew.process.jsonl_adapter import SQLJsonlConverter -from datastew.repository.sqllite import SQLLiteRepository -from tests.base_repository_test_setup import BaseRepositoryTestSetup - - -class TestSQLLiteRepository(BaseRepositoryTestSetup): - __test__ = True - - @classmethod - def setUpClass(cls): - super().setUpClass() - cls.repo_args = ("disk", "sqlite_db", cls.vectorizer1) - cls.repository = SQLLiteRepository(*cls.repo_args) - cls.jsonl_converter = SQLJsonlConverter(dest_dir="test_export")