diff --git a/datastew/process/jsonl_adapter.py b/datastew/process/jsonl_adapter.py index ca6614f..86d3964 100644 --- a/datastew/process/jsonl_adapter.py +++ b/datastew/process/jsonl_adapter.py @@ -1,12 +1,17 @@ import json import os +from abc import ABC, abstractmethod +from typing import Any, Dict, Union +import numpy as np import pandas as pd from tqdm import tqdm from weaviate.util import generate_uuid5 from datastew.embedding import Vectorizer -from datastew.repository import WeaviateRepository +from datastew.repository import PostgreSQLRepository, SQLLiteRepository, WeaviateRepository +from datastew.repository.base import BaseRepository +from datastew.repository.model import Concept, Mapping, Terminology from datastew.repository.weaviate_schema import ( concept_schema, mapping_schema_user_vectors, @@ -14,33 +19,20 @@ ) -class WeaviateJsonlConverter(object): - """ - Converts data to our JSONL format for Weaviate schema. - """ +class BaseJsonlConverter(ABC): + def __init__(self, dest_dir: str, buffer_size: int = 1000): + """Initialize the converter. - def __init__( - self, - dest_dir: str, - terminology_schema: dict = terminology_schema.schema, - concept_schema: dict = concept_schema.schema, - mapping_schema: dict = mapping_schema_user_vectors.schema, - buffer_size: int = 1000, - ): + :param dest_dir: Destination directory for the exported JSONL files. + :param buffer_size: Number of records to buffer before writing to disk, defaults to 1000 + """ self.dest_dir = dest_dir - self.terminology_schema = terminology_schema - self.concept_schema = concept_schema - self.mapping_schema = mapping_schema self._buffer = [] self._buffer_size = buffer_size self._ensure_directories_exist() def _ensure_directories_exist(self): - """ - Ensures the output directory exists. - - :return: None - """ + """Ensures the output directory exists.""" os.makedirs(self.dest_dir, exist_ok=True) def _get_file_path(self, collection: str) -> str: @@ -79,33 +71,187 @@ def _flush_to_file(self, file_path: str): with open(file_path, "a", encoding="utf-8") as file: for entry in self._buffer: + if isinstance(entry.get("embedding"), np.ndarray): + entry["embedding"] = entry["embedding"].tolist() file.write(json.dumps(entry) + "\n") self._buffer.clear() - def from_repository(self, repository: WeaviateRepository) -> None: + @abstractmethod + def from_repository(self, repository: BaseRepository): + pass + + @abstractmethod + def from_ohdsi(self, src: str, vectorizer: Vectorizer = Vectorizer(), include_vectors: bool = True): + pass + + @abstractmethod + def _object_to_dict(self, obj: Any) -> Dict[str, Any]: + pass + + +class SQLJsonlConverter(BaseJsonlConverter): + def __init__(self, dest_dir: str, buffer_size: int = 1000): + super().__init__(dest_dir=dest_dir, buffer_size=buffer_size) + + def from_repository(self, repository: Union[PostgreSQLRepository, SQLLiteRepository]): + """Export all records from a PostgreSQLRepository to JSONL files + + :param repository: Active database repository instace. + """ + session = repository.session + + # Export Terminologies + terminology_file_path = self._get_file_path("terminology") + for t in tqdm(session.query(Terminology).all(), desc="Exporting Terminologies"): + terminology = self._object_to_dict(t) + self._write_to_jsonl(terminology_file_path, terminology) + self._flush_to_file(terminology_file_path) + + # Export Concepts + concept_file_path = self._get_file_path("concept") + for c in tqdm(session.query(Concept).all(), desc="Exporting Concepts"): + concept = self._object_to_dict(c) + self._write_to_jsonl(concept_file_path, concept) + self._flush_to_file(concept_file_path) + + # Export Mappings + mapping_file_path = self._get_file_path("mapping") + for m in tqdm(session.query(Mapping).all(), desc="Exporting Mappings"): + mapping = self._object_to_dict(m) + self._write_to_jsonl(mapping_file_path, mapping) + self._flush_to_file(mapping_file_path) + + def from_ohdsi(self, src: str, vectorizer: Vectorizer = Vectorizer(), include_vectors: bool = True): + """ + Converts data from OHDSI to SQL-compatible JSONL format. + + :param src: Path to the OHDSI CONCEPT.csv file. + :param vectorizer: Vectorizer to use for text embeddings. + :param include_vectors: Whether to include vector data in mappings. + """ + if not os.path.exists(src): + raise FileNotFoundError(f"OHDSI concept file '{src}' does not exist or is not a file.") + + terminology_file_path = self._get_file_path("terminology") + concept_file_path = self._get_file_path("concept") + mapping_file_path = self._get_file_path("mapping") + + # Write single OHDSI terminology entry + self._write_to_jsonl(terminology_file_path, {"id": "OHDSI", "name": "OHDSI"}) + self._flush_to_file(terminology_file_path) + + for chunk in tqdm( + pd.read_csv( + src, + delimiter="\t", + usecols=["concept_name", "concept_id"], + chunksize=10000, + dtype={"concept_id": str, "concept_name": str}, + ), + desc="Processing OHDSI concepts", + ): + + concepts = [] + mappings = [] + + concept_names = chunk["concept_name"].astype(str).tolist() + concept_ids = chunk["concept_id"].astype(str).tolist() + + if include_vectors: + embeddings = vectorizer.get_embeddings(concept_names) + + for i in range(len(concept_names)): + concept_identifier = f"OHDSI:{concept_ids[i]}" + label = concept_names[i] + + # Concept JSON + concepts.append( + { + "concept_identifier": concept_identifier, + "pref_label": label, + "terminology_id": "OHDSI", + } + ) + + # Mapping JSON + mapping = { + "concept_identifier": concept_identifier, + "text": label, + } + if include_vectors: + mapping["sentence_embedder"] = vectorizer.model_name + mapping["embedding"] = embeddings[i] + + mappings.append(mapping) + + # Write results in batch + for concept_data in concepts: + self._write_to_jsonl(concept_file_path, concept_data) + self._flush_to_file(concept_file_path) + for mapping_data in mappings: + self._write_to_jsonl(mapping_file_path, mapping_data) + self._flush_to_file(mapping_file_path) + + def _object_to_dict(self, obj: Union[Terminology, Concept, Mapping]) -> Dict[str, Any]: + if isinstance(obj, Terminology): + return { + "id": obj.id, + "name": obj.name, + } + elif isinstance(obj, Concept): + return { + "concept_identifier": obj.concept_identifier, + "pref_label": obj.pref_label, + "terminology_id": obj.terminology.id, + } + elif isinstance(obj, Mapping): + return { + "concept_identifier": obj.concept.concept_identifier, + "text": obj.text, + "embedding": obj.embedding, + "sentence_embedder": obj.sentence_embedder, + } + else: + raise TypeError(f"Unsupported object type: {type(obj)}") + + +class WeaviateJsonlConverter(BaseJsonlConverter): + def __init__( + self, + dest_dir: str, + terminology_schema: dict = terminology_schema.schema, + concept_schema: dict = concept_schema.schema, + mapping_schema: dict = mapping_schema_user_vectors.schema, + buffer_size: int = 1000, + ): + super().__init__(dest_dir=dest_dir, buffer_size=buffer_size) + self.terminology_schema = terminology_schema + self.concept_schema = concept_schema + self.mapping_schema = mapping_schema + + def from_repository(self, repository: WeaviateRepository): """ Converts data from a WeaviateRepository to our JSONL format. :param repository: WeaviateRepository - :return: None """ # Process terminology first terminology_file_path = self._get_file_path("terminology") for terminology in repository.get_iterator(self.terminology_schema["class"]): - self._write_to_jsonl(terminology_file_path, self._weaviate_object_to_dict(terminology)) + self._write_to_jsonl(terminology_file_path, self._object_to_dict(terminology)) self._flush_to_file(terminology_file_path) # Process concept next concept_file_path = self._get_file_path("concept") for concept in repository.get_iterator(self.concept_schema["class"]): - self._write_to_jsonl(concept_file_path, self._weaviate_object_to_dict(concept)) + self._write_to_jsonl(concept_file_path, self._object_to_dict(concept)) self._flush_to_file(concept_file_path) # Process mapping last mapping_file_path = self._get_file_path("mapping") for mapping in repository.get_iterator(self.mapping_schema["class"]): - self._write_to_jsonl(mapping_file_path, self._weaviate_object_to_dict(mapping)) + self._write_to_jsonl(mapping_file_path, self._object_to_dict(mapping)) self._flush_to_file(mapping_file_path) def from_ohdsi(self, src: str, vectorizer: Vectorizer = Vectorizer(), include_vectors: bool = True): @@ -128,7 +274,7 @@ def from_ohdsi(self, src: str, vectorizer: Vectorizer = Vectorizer(), include_ve terminology_properties = {"name": "OHDSI"} terminology_id = generate_uuid5(terminology_properties) ohdsi_terminology = { - "class": self.terminology_schema["class"], + "class": "Terminology", "id": terminology_id, "properties": terminology_properties, } @@ -206,22 +352,25 @@ def from_ohdsi(self, src: str, vectorizer: Vectorizer = Vectorizer(), include_ve self._write_to_jsonl(mapping_file_path, mapping_data) self._flush_to_file(mapping_file_path) - @staticmethod - def _weaviate_object_to_dict(weaviate_object): + def _object_to_dict(self, obj) -> Dict[str, Any]: + """Conver a Weaviate object to a schema-compliant JSON dictionary. - if weaviate_object.references is not None: + :param obj: Weaviate object instance. + :return: Formatted dictionary containing class, id, properties, vector, and references. + """ + if obj.references is not None: # FIXME: This is a hack to get the UUID of the referenced object. Replace as soon as weaviate devs offer an # actual solution for this. - vals = [value.objects for _, value in weaviate_object.references.items()] + vals = [value.objects for _, value in obj.references.items()] uuid = [str(obj.uuid) for sublist in vals for obj in sublist][0] - references = {key: uuid for key, _ in weaviate_object.references.items()} + references = {key: uuid for key, _ in obj.references.items()} else: references = {} return { - "class": weaviate_object.collection, - "id": str(weaviate_object.uuid), - "properties": weaviate_object.properties, - "vector": weaviate_object.vector, + "class": obj.collection, + "id": str(obj.uuid), + "properties": obj.properties, + "vector": obj.vector, "references": references, } diff --git a/datastew/repository/__init__.py b/datastew/repository/__init__.py index 8635a29..7deea5e 100644 --- a/datastew/repository/__init__.py +++ b/datastew/repository/__init__.py @@ -1,14 +1,6 @@ from .model import Concept, Mapping, Terminology +from .postgresql import PostgreSQLRepository from .sqllite import SQLLiteRepository from .weaviate import WeaviateRepository -from .weaviate_schema import (concept_schema, - mapping_schema_preconfigured_embeddings, - mapping_schema_user_vectors, terminology_schema) -__all__ = [ - "Terminology", - "Concept", - "Mapping", - "SQLLiteRepository", - "WeaviateRepository", -] +__all__ = ["Terminology", "Concept", "Mapping", "SQLLiteRepository", "WeaviateRepository", "PostgreSQLRepository"] diff --git a/datastew/repository/base.py b/datastew/repository/base.py index fe89fd0..d8138dd 100644 --- a/datastew/repository/base.py +++ b/datastew/repository/base.py @@ -15,12 +15,12 @@ def __init__(self, vectorizer: Vectorizer = Vectorizer()): self.vectorizer = vectorizer @abstractmethod - def store(self, model_object_instance): + 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): + def store_all(self, model_object_instances: List[Union[Terminology, Concept, Mapping]]): """Store multiple model object instances.""" pass @@ -30,7 +30,7 @@ def get_concept(self, concept_id: str) -> Concept: pass @abstractmethod - def get_concepts(self) -> Page[Concept]: + def get_concepts(self, terminology_name: Optional[str] = None, offset: int = 0, limit: int = 100) -> Page[Concept]: """Retrieve all concepts from the database.""" pass @@ -95,6 +95,10 @@ def import_data_dictionary(self, data_dictionary: DataDictionarySource, terminol 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]]: diff --git a/datastew/repository/model.py b/datastew/repository/model.py index c809a33..e991cbe 100644 --- a/datastew/repository/model.py +++ b/datastew/repository/model.py @@ -73,7 +73,6 @@ class Concept(Base): pref_label = Column(String) terminology_id = Column(String, ForeignKey("terminology.id")) terminology = relationship("Terminology") - uuid = Column(String) def __init__(self, terminology: Terminology, pref_label: str, concept_identifier: str, id: Optional[str] = None): self.terminology = terminology diff --git a/datastew/repository/postgresql.py b/datastew/repository/postgresql.py index 82e95b5..d71e353 100644 --- a/datastew/repository/postgresql.py +++ b/datastew/repository/postgresql.py @@ -1,5 +1,6 @@ +import json import logging -from typing import List, Optional, Sequence, Union +from typing import Any, Dict, List, Literal, Optional, Sequence, Union from sqlalchemy import create_engine, func, inspect, text from sqlalchemy.orm import joinedload, sessionmaker @@ -199,6 +200,111 @@ def clear_all(self): 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}") + def _initialize_pgvector(self): with self.engine.begin() as conn: conn.execute(text("CREATE EXTENSION IF NOT EXISTS vector")) diff --git a/datastew/repository/sqllite.py b/datastew/repository/sqllite.py index 954c793..5bb779f 100644 --- a/datastew/repository/sqllite.py +++ b/datastew/repository/sqllite.py @@ -1,5 +1,6 @@ +import json import logging -from typing import List, Optional, Union +from typing import Any, Dict, List, Literal, Optional, Union import numpy as np from sqlalchemy import create_engine, func, inspect @@ -216,6 +217,111 @@ def clear_all(self): 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}") + def _is_duplicate(self, model_object_instance: Union[Terminology, Concept, Mapping]) -> bool: """Checks whether an object with the same primary key already exists. diff --git a/datastew/repository/weaviate.py b/datastew/repository/weaviate.py index a7d4432..0eebfb0 100644 --- a/datastew/repository/weaviate.py +++ b/datastew/repository/weaviate.py @@ -1,6 +1,5 @@ import json import logging -import shutil import socket from typing import Any, Dict, List, Literal, Optional, Sequence, Union @@ -671,11 +670,13 @@ def store(self, model_object_instance: Union[Terminology, Concept, Mapping]): except Exception as e: raise ObjectStorageError("Failed to store object in the database.", e) - def import_from_jsonl(self, jsonl_path: str, object_type: str, chunk_size: int = 100): + 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 Weaviate database. :param jsonl_path: Path to the JSONL file. - :param object_type: The type of objects to import ("terminology", "concept", "mapping"). + :param object_type: Literal specifying the object type, must be "terminology", "concept", or "mapping". :param chunk_size: The number of items to process in each batch, defaults to 100. :raises ValueError: If the client is not initialized or is invalid. :raises ValueError: If 'id' or 'properties' is missing in a JSON object. @@ -717,17 +718,15 @@ def import_from_jsonl(self, jsonl_path: str, object_type: str, chunk_size: int = except Exception as e: raise RuntimeError(f"An unexpected error occurred during import: {e}") + @deprecated("close is deprecated and will be removed in a future release, use shut_down instead") def close(self): + self.shut_down() + + def shut_down(self): if not self.client: raise ValueError("Client is not initialized or is invalid.") self.client.close() - def shut_down(self): - if self.mode == "memory": - shutil.rmtree("db") - else: - self.close() - def clear_all(self): """Deletes all data and schema classes (Mapping, Concept, Terminology) and re-creates them. diff --git a/tests/base_repository_test_setup.py b/tests/base_repository_test_setup.py index c4d3b9d..f9d1b38 100644 --- a/tests/base_repository_test_setup.py +++ b/tests/base_repository_test_setup.py @@ -3,6 +3,7 @@ from typing import List, Tuple from datastew.embedding import Vectorizer +from datastew.process.jsonl_adapter import BaseJsonlConverter from datastew.process.parsing import DataDictionarySource from datastew.repository import Concept, Mapping, Terminology from datastew.repository.base import BaseRepository @@ -10,8 +11,10 @@ class BaseRepositoryTestSetup(unittest.TestCase): """Base class for setting up test data and shared tests for all repository backends.""" + __test__ = False repository: BaseRepository + jsonl_converter: BaseJsonlConverter TEST_CONCEPTS = [ ("Diabetes mellitus (disorder)", "Concept ID: 11893007", "v1"), @@ -171,3 +174,29 @@ def test_repository_restart(self): concepts = repo.get_concepts() self.assertEqual(len(concepts.items), 11) + + def test_jsonl_export(self): + converter = self.jsonl_converter + converter.from_repository(self.repository) + # assert that the dest dir + self.assertTrue(converter.dest_dir) + # assert that the files is not empty + with open(converter.dest_dir + "/terminology.jsonl", "r") as file: + self.assertTrue(file.read()) + with open(converter.dest_dir + "/concept.jsonl", "r") as file: + self.assertTrue(file.read()) + with open(converter.dest_dir + "/mapping.jsonl", "r") as file: + self.assertTrue(file.read()) + # assert that the file contains the expected data + with open(converter.dest_dir + "/terminology.jsonl", "r") as file: + self.assertIn("snomed CT", file.read()) + with open(converter.dest_dir + "/concept.jsonl", "r") as file: + self.assertIn("Diabetes mellitus (disorder)", file.read()) + with open(converter.dest_dir + "/mapping.jsonl", "r") as file: + self.assertIn("Diabetes mellitus (disorder)", file.read()) + # remove the created dir and files + os.remove(converter.dest_dir + "/terminology.jsonl") + os.remove(converter.dest_dir + "/concept.jsonl") + os.remove(converter.dest_dir + "/mapping.jsonl") + # remove the created dir + os.rmdir(converter.dest_dir) diff --git a/tests/test_postgreqsl_repository.py b/tests/test_postgreqsl_repository.py index 03466e4..aee0633 100644 --- a/tests/test_postgreqsl_repository.py +++ b/tests/test_postgreqsl_repository.py @@ -1,5 +1,6 @@ import os +from datastew.process.jsonl_adapter import SQLJsonlConverter from datastew.repository.postgresql import PostgreSQLRepository from tests.base_repository_test_setup import BaseRepositoryTestSetup @@ -13,3 +14,4 @@ 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/test_postgresql_import.py b/tests/test_postgresql_import.py new file mode 100644 index 0000000..60cd757 --- /dev/null +++ b/tests/test_postgresql_import.py @@ -0,0 +1,126 @@ +import json +import os +import random +import shutil +import tempfile +from typing import Any, Dict, List, Literal +from unittest import TestCase + +from datastew.repository.postgresql import PostgreSQLRepository + + +class TestWeaviateRepositoryImport(TestCase): + + def setUp(self) -> None: + POSTGRES_TEST_URL = os.getenv("TEST_POSTGRES_URI", "postgresql://testuser:testpass@localhost/testdb") + self.repository = PostgreSQLRepository(POSTGRES_TEST_URL) + 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": [random.uniform(-1, 1) for _ in range(768)], + "sentence_embedder": "sentence-transformers/all-mpnet-base-v2", + }, + { + "text": "liver", + "concept_identifier": "import_test:H", + "embedding": [random.uniform(-1, 1) for _ in range(768)], + "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), 768) + 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_import.py b/tests/test_sqllite_import.py new file mode 100644 index 0000000..e131309 --- /dev/null +++ b/tests/test_sqllite_import.py @@ -0,0 +1,124 @@ +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 index 018ea5f..0ef94f4 100644 --- a/tests/test_sqllite_repository.py +++ b/tests/test_sqllite_repository.py @@ -1,3 +1,4 @@ +from datastew.process.jsonl_adapter import SQLJsonlConverter from datastew.repository.sqllite import SQLLiteRepository from tests.base_repository_test_setup import BaseRepositoryTestSetup @@ -10,3 +11,4 @@ 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") diff --git a/tests/test_system.py b/tests/test_system.py deleted file mode 100644 index 62fb1dd..0000000 --- a/tests/test_system.py +++ /dev/null @@ -1,61 +0,0 @@ -import unittest - -from datastew.embedding import Vectorizer -from datastew.repository.model import Concept, Mapping, Terminology -from datastew.repository.sqllite import SQLLiteRepository - - -class TestGetClosestEmbedding(unittest.TestCase): - - def setUp(self): - self.repository = SQLLiteRepository(mode="memory") - self.vectorizer = Vectorizer() - - def tearDown(self): - self.repository.shut_down() - - def test_mapping_storage_and_closest_retrieval(self): - # preset knowledge - terminology = Terminology("test", "test") - concept1 = Concept(terminology, "cat", "TEST:1") - concept1_description = "The cat is sitting on the mat." - sentence_embedder = "test" - mapping1 = Mapping( - concept1, - concept1_description, - list(self.vectorizer.get_embedding(concept1_description)), - sentence_embedder=sentence_embedder, - ) - concept2 = Concept(terminology, "sunrise", "TEST:2") - concept2_description = "The sun rises in the east." - mapping2 = Mapping( - concept2, - concept2_description, - list(self.vectorizer.get_embedding(concept2_description)), - sentence_embedder=sentence_embedder, - ) - concept3 = Concept(terminology, "dog", "TEST:3") - concept3_description = "A loyal companion to humans." - mapping3 = Mapping( - concept3, - concept3_description, - list(self.vectorizer.get_embedding(concept3_description)), - sentence_embedder=sentence_embedder, - ) - self.repository.store_all([terminology, concept1, mapping1, concept2, mapping2, concept3, mapping3]) - # test new mappings - text1 = "A furry feline rests on the rug." - text1_embedding = self.vectorizer.get_embedding(text1) - text2 = "Dawn breaks over the horizon." - text2_embedding = self.vectorizer.get_embedding(text2) - text3 = "A faithful friend." - text3_embedding = self.vectorizer.get_embedding(text3) - mappings1 = self.repository.get_closest_mappings(list(text1_embedding), limit=3) - mappings2 = self.repository.get_closest_mappings(list(text2_embedding), limit=3) - mappings3 = self.repository.get_closest_mappings(list(text3_embedding), limit=3) - self.assertEqual(len(mappings1), 3) - self.assertEqual(len(mappings2), 3) - self.assertEqual(len(mappings3), 3) - self.assertEqual(concept1_description, mappings1[0].mapping.text) - self.assertEqual(concept2_description, mappings2[0].mapping.text) - self.assertEqual(concept3_description, mappings3[0].mapping.text) diff --git a/tests/test_weaviate_export.py b/tests/test_weaviate_export.py deleted file mode 100644 index 1b1d232..0000000 --- a/tests/test_weaviate_export.py +++ /dev/null @@ -1,58 +0,0 @@ -import os -import shutil -from unittest import TestCase - -from datastew import Concept, Mapping, Terminology -from datastew.embedding import Vectorizer -from datastew.process.jsonl_adapter import WeaviateJsonlConverter -from datastew.repository import WeaviateRepository - - -class TestWeaviateRepositoryExport(TestCase): - - @classmethod - def setUp(cls) -> None: - - cls.repository = WeaviateRepository() - terminology = Terminology("snomed CT", "SNOMED") - - vectorizer = Vectorizer() - - text1 = "Diabetes mellitus (disorder)" - concept1 = Concept(terminology, text1, "Concept ID: 11893007") - mapping1 = Mapping(concept1, text1, vectorizer.get_embedding(text1), vectorizer.model_name) - - cls.repository.store_all([terminology, concept1, mapping1]) - - @classmethod - def tearDownClass(cls) -> None: - cls.repository.close() - shutil.rmtree(os.path.join(os.getcwd(), "db")) - - def test_jsonl_export(self): - converter = WeaviateJsonlConverter(dest_dir="test_export") - converter.from_repository(self.repository) - # assert that the dest dir - self.assertTrue(converter.dest_dir) - # assert that the files is not empty - with open(converter.dest_dir + "/terminology.jsonl", 'r') as file: - self.assertTrue(file.read()) - with open(converter.dest_dir + "/concept.jsonl", 'r') as file: - self.assertTrue(file.read()) - with open(converter.dest_dir + "/mapping.jsonl", 'r') as file: - self.assertTrue(file.read()) - # assert that the file contains the expected data - with open(converter.dest_dir + "/terminology.jsonl", 'r') as file: - self.assertIn("snomed CT", file.read()) - with open(converter.dest_dir + "/concept.jsonl", 'r') as file: - self.assertIn("Diabetes mellitus (disorder)", file.read()) - with open(converter.dest_dir + "/mapping.jsonl", 'r') as file: - self.assertIn("Diabetes mellitus (disorder)", file.read()) - # remove the created dir and files - os.remove(converter.dest_dir + "/terminology.jsonl") - os.remove(converter.dest_dir + "/concept.jsonl") - os.remove(converter.dest_dir + "/mapping.jsonl") - # remove the created dir - os.rmdir(converter.dest_dir) - # close the db connection - self.repository.close() \ No newline at end of file diff --git a/tests/test_weaviate_import.py b/tests/test_weaviate_import.py index 343d5d6..0418ad9 100644 --- a/tests/test_weaviate_import.py +++ b/tests/test_weaviate_import.py @@ -2,7 +2,7 @@ import os import shutil import tempfile -from typing import Any, Dict, List +from typing import Any, Dict, List, Literal from unittest import TestCase from datastew.repository import WeaviateRepository @@ -12,6 +12,7 @@ class TestWeaviateRepositoryImport(TestCase): def setUp(self) -> None: self.repository = WeaviateRepository() + self.repository.clear_all() self.temp_dir = tempfile.mkdtemp() # Sample data for JSONL files @@ -31,18 +32,14 @@ def setUp(self) -> None: "id": "064cb594-41cd-561d-b5a8-2bf226006f09", "properties": {"conceptID": "import_test:G", "prefLabel": "G"}, "vector": {}, - "references": { - "hasTerminology": "94331523-fa7e-5871-9375-8f559d6035dd" - }, + "references": {"hasTerminology": "94331523-fa7e-5871-9375-8f559d6035dd"}, }, { "class": "Concept", "id": "12345678-41cd-561d-b5a8-2bf226006f09", "properties": {"conceptID": "import_test:H", "prefLabel": "H"}, "vector": {}, - "references": { - "hasTerminology": "94331523-fa7e-5871-9375-8f559d6035dd" - }, + "references": {"hasTerminology": "94331523-fa7e-5871-9375-8f559d6035dd"}, }, ], "mapping": [ @@ -53,12 +50,8 @@ def setUp(self) -> None: "text": "pancreas", "hasSentenceEmbedder": "sentence-transformers/all-mpnet-base-v2", }, - "vector": { - "default": [0.048744574189186096, -0.0035385489463806152] - }, - "references": { - "hasConcept": "064cb594-41cd-561d-b5a8-2bf226006f09" - }, + "vector": {"default": [0.048744574189186096, -0.0035385489463806152]}, + "references": {"hasConcept": "064cb594-41cd-561d-b5a8-2bf226006f09"}, }, { "class": "Mapping", @@ -68,9 +61,7 @@ def setUp(self) -> None: "hasSentenceEmbedder": "sentence-transformers/all-mpnet-base-v2", }, "vector": {"default": [0.1, -0.2]}, - "references": { - "hasConcept": "12345678-41cd-561d-b5a8-2bf226006f09" - }, + "references": {"hasConcept": "12345678-41cd-561d-b5a8-2bf226006f09"}, }, ], } @@ -87,7 +78,7 @@ def write_jsonl(file_path: str, data: List[Dict[str, Any]]): json.dump(obj, file) file.write("\n") - def import_data(self, data_types: List[str]): + 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") @@ -95,8 +86,7 @@ def import_data(self, data_types: List[str]): def tearDown(self) -> None: shutil.rmtree(self.temp_dir, ignore_errors=True) - self.repository.close() - shutil.rmtree(os.path.join(os.getcwd(), "db"), ignore_errors=True) + self.repository.shut_down() def test_import_terminology(self): self.import_data(["terminology"]) @@ -125,9 +115,7 @@ def test_import_concepts(self): ) 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"] - ) + 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") @@ -149,9 +137,7 @@ def test_import_mappings(self): 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" - ) + 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}"): @@ -168,9 +154,7 @@ def test_import_invalid_jsonl(self): def test_import_missing_id(self): file_path = os.path.join(self.temp_dir, "missing_id.jsonl") - self.write_jsonl( - file_path, [{"class": "Terminology", "properties": {"name": "missing_id"}}] - ) + self.write_jsonl(file_path, [{"class": "Terminology", "properties": {"name": "missing_id"}}]) with self.assertRaises(ValueError): self.repository.import_from_jsonl(file_path, "terminology") diff --git a/tests/test_weaviate_repository.py b/tests/test_weaviate_repository.py index 3018610..25d47d6 100644 --- a/tests/test_weaviate_repository.py +++ b/tests/test_weaviate_repository.py @@ -1,3 +1,4 @@ +from datastew.process.jsonl_adapter import WeaviateJsonlConverter from datastew.repository.weaviate import WeaviateRepository from tests.base_repository_test_setup import BaseRepositoryTestSetup @@ -9,3 +10,4 @@ class TestWeaviateRepository(BaseRepositoryTestSetup): def setUpClass(cls): super().setUpClass() cls.repository = WeaviateRepository(vectorizer=cls.vectorizer1) + cls.jsonl_converter = WeaviateJsonlConverter(dest_dir="test_export")