diff --git a/README.md b/README.md index 4d99898..b6a4e0c 100644 --- a/README.md +++ b/README.md @@ -14,7 +14,7 @@ dependency. ## Installation ```bash -pip install dataconnect # core (pyarrow + pydantic + httpx) +pip install dataconnect # core (pyarrow) pip install dataconnect[pandas] # + pandas for .to_pandas() on results ``` @@ -25,7 +25,8 @@ Requires **Python ≥ 3.13**. ## Quick start ```python -import pyarrow as pa +from uuid import UUID + from dataconnect import DataConnectClient @@ -36,6 +37,9 @@ with DataConnectClient.connect( ) as client: studies = client.get_studies(search_study_name="ACME") + + datasets = client.get_datasets(study_environment_uuid=UUID("cec9f2a7-07ba-4fa8-bfcf-34fbc5d56793")) + ``` ## Development diff --git a/dataconnect/client.py b/dataconnect/client.py index 9473a30..83e7ef1 100644 --- a/dataconnect/client.py +++ b/dataconnect/client.py @@ -10,7 +10,7 @@ from types import TracebackType from uuid import UUID -from dataconnect.models import DatasetVersion, Study +from dataconnect.models import Dataset, DatasetVersion, Study from dataconnect.service import DataConnectService, DefaultDataConnectService _DEFAULT_HOST = "enodia-gateway.platform.imedidata.com" @@ -51,6 +51,31 @@ def get_dataset_versions(self, dataset_uuid: UUID) -> list[DatasetVersion]: """List the dataset versions the client is authorized to access.""" return self._service.get_dataset_versions(dataset_uuid) + def get_datasets( + self, + study_environment_uuid: UUID, + search_dataset_name: str = "", + page: int = 1, + page_size: int = 50, + ) -> list[Dataset]: + """List datasets for a study environment. + + Args: + study_environment_uuid: UUID of the study environment (required). + search_dataset_name: Full or partial dataset name filter. + page: Page number for paginated results. + page_size: Number of results per page. + + Returns: + A list of :class:`Dataset` items matching the criteria. + """ + return self._service.get_datasets( + study_environment_uuid=study_environment_uuid, + search_dataset_name=search_dataset_name, + page=page, + page_size=page_size, + ) + # Lifecycle def close(self) -> None: diff --git a/dataconnect/models.py b/dataconnect/models.py index 117f560..cc713b7 100644 --- a/dataconnect/models.py +++ b/dataconnect/models.py @@ -1,8 +1,11 @@ from __future__ import annotations from dataclasses import dataclass, field +from typing import Generic, TypeVar from uuid import UUID +T = TypeVar("T") + @dataclass(frozen=True) class StudyEnvironment: @@ -24,3 +27,31 @@ class DatasetVersion: dataset_uuid: UUID dataset_name: str dataset_version: str + + +@dataclass(frozen=True) +class Dataset: + """A dataset belonging to a study environment.""" + + dataset_uuid: str + study_uuid: str + study_env_uuid: str + dataset_name: str + + +@dataclass(frozen=True) +class Pagination: + """Server-side pagination metadata.""" + + page: int + page_size: int + total_pages: int + + +@dataclass +class PaginatedResponse(Generic[T]): # noqa: UP046 + """A paginated collection returned by list endpoints.""" + + total_records: int + pagination: Pagination + items: list[T] diff --git a/dataconnect/service/base.py b/dataconnect/service/base.py index 338d340..ab0a224 100644 --- a/dataconnect/service/base.py +++ b/dataconnect/service/base.py @@ -5,7 +5,7 @@ from abc import ABC, abstractmethod from uuid import UUID -from dataconnect.models import DatasetVersion, Study +from dataconnect.models import Dataset, DatasetVersion, Study class DataConnectService(ABC): @@ -14,6 +14,15 @@ class DataConnectService(ABC): @abstractmethod def get_studies(self, search_study_name: str | None = None) -> list[Study]: ... + @abstractmethod + def get_datasets( + self, + study_environment_uuid: UUID, + search_dataset_name: str = "", + page: int = 1, + page_size: int = 50, + ) -> list[Dataset]: ... + @abstractmethod def get_dataset_versions(self, dataset_uuid: UUID) -> list[DatasetVersion]: ... diff --git a/dataconnect/service/default.py b/dataconnect/service/default.py index 5272ea8..c67fac3 100644 --- a/dataconnect/service/default.py +++ b/dataconnect/service/default.py @@ -14,9 +14,9 @@ ServerError, ValidationError, ) -from dataconnect.models import DatasetVersion, Study +from dataconnect.models import Dataset, DatasetVersion, Study from dataconnect.service.base import DataConnectService -from dataconnect.service.mappers import resource_to_dataset_version, resource_to_study +from dataconnect.service.mappers import resource_to_dataset, resource_to_dataset_version, resource_to_study from dataconnect.service.validators import validate_search_study_name from dataconnect.transport.base import Transport from dataconnect.transport.errors import ( @@ -32,6 +32,7 @@ # Server action identifiers _ACTION_LIST_STUDIES = "studies.list" +_ACTION_LIST_DATASETS = "datasets.list" _ACTION_LIST_DATASET_VERSIONS = "dataset_versions.list" @@ -100,6 +101,49 @@ def get_dataset_versions(self, dataset_uuid: UUID) -> list[DatasetVersion]: except (IndexError, KeyError, TypeError, ValueError) as ex: raise ValidationError(f"Unexpected dataset versions response format: {ex}") from ex + def get_datasets( + self, + study_environment_uuid: UUID, + search_dataset_name: str = "", + page: int = 1, + page_size: int = 50, + ) -> list[Dataset]: + """List datasets for a study environment. + + Args: + study_environment_uuid: UUID of the study environment (required). + search_dataset_name: Full or partial dataset name filter. + page: Page number for paginated results. + page_size: Number of results per page. + + Returns: + A list of :class:`Dataset` items matching the criteria. + """ + if not isinstance(study_environment_uuid, UUID): + raise ValidationError("study_environment_uuid must be a valid UUID") + + if study_environment_uuid.int == 0: + raise ValidationError("study_environment_uuid must not be empty") + + request = ResourceQuery(action=_ACTION_LIST_DATASETS).append_body( + { + "study_environment_uuid": str(study_environment_uuid), + "search_dataset_name": search_dataset_name, + "page": page, + "page_size": page_size, + } + ) + + try: + resources = self._transport.list_resources(request) + except TransportError as ex: + raise _translate_error(ex) from ex + + try: + return [resource_to_dataset(r) for r in resources] + except (IndexError, KeyError, TypeError, ValueError) as ex: + raise ValidationError(f"Unexpected datasets response format: {ex}") from ex + def close(self) -> None: try: diff --git a/dataconnect/service/mappers.py b/dataconnect/service/mappers.py index 1422b81..00d93d2 100644 --- a/dataconnect/service/mappers.py +++ b/dataconnect/service/mappers.py @@ -11,7 +11,7 @@ from uuid import UUID from dataconnect.exceptions import NotFoundError -from dataconnect.models import DatasetVersion, Study, StudyEnvironment +from dataconnect.models import Dataset, DatasetVersion, Study, StudyEnvironment from dataconnect.transport.models import ResourceInfo @@ -45,3 +45,19 @@ def resource_to_dataset_version(resource: ResourceInfo) -> DatasetVersion: dataset_name=data["dataset_name"], dataset_version=data["dataset_version"], ) + + +def resource_to_dataset(resource: ResourceInfo) -> Dataset: + """Parse a transport-layer ``ResourceInfo`` into a ``Dataset`` domain object.""" + + if not resource or not resource.endpoints or not resource.endpoints[0].ticket: + raise NotFoundError("Invalid resource: missing endpoints or ticket") + + data = json.loads(resource.endpoints[0].ticket.decode("utf-8")) + + return Dataset( + dataset_uuid=data.get("dataset_uuid", ""), + study_uuid=data.get("study_uuid", ""), + study_env_uuid=data.get("study_env_uuid", ""), + dataset_name=data.get("dataset_name", ""), + ) diff --git a/dataconnect/transport/arrow_flight/transport.py b/dataconnect/transport/arrow_flight/transport.py index dabd846..824163b 100644 --- a/dataconnect/transport/arrow_flight/transport.py +++ b/dataconnect/transport/arrow_flight/transport.py @@ -42,6 +42,7 @@ def _to_resource_info(info: flight.FlightInfo) -> ResourceInfo: # Maps service-layer action names to the flight_type value the Arrow Flight server expects. _ACTION_FLIGHT_TYPE: dict[str, str] = { "studies.list": "STUDIES", + "datasets.list": "DATASETS", "dataset_versions.list": "VERSIONS", } diff --git a/tests/test_client.py b/tests/test_client.py index 92ad150..760bd20 100644 --- a/tests/test_client.py +++ b/tests/test_client.py @@ -7,15 +7,22 @@ import pytest from dataconnect.client import DataConnectClient -from dataconnect.models import DatasetVersion, Study +from dataconnect.models import Dataset, DatasetVersion, Study class _FakeService: - def __init__(self, studies: list[Study] | None = None, versions: list[DatasetVersion] | None = None) -> None: + def __init__( + self, + studies: list[Study] | None = None, + versions: list[DatasetVersion] | None = None, + datasets: list[Dataset] | None = None, + ) -> None: self._studies = studies or [] self._versions = versions or [] + self._datasets = datasets or [] self.closed = 0 self.last_dataset_uuid: UUID | None = None + self.last_get_datasets_kwargs: dict[str, object] | None = None def get_studies(self) -> list[Study]: return self._studies @@ -24,6 +31,10 @@ def get_dataset_versions(self, dataset_uuid: UUID) -> list[DatasetVersion]: self.last_dataset_uuid = dataset_uuid return self._versions + def get_datasets(self, **kwargs: object) -> list[Dataset]: + self.last_get_datasets_kwargs = kwargs + return self._datasets + def close(self) -> None: self.closed += 1 @@ -144,3 +155,46 @@ def test_close_delegates_to_service() -> None: client.close() assert service.was_closed + + +def test_get_datasets_forwards_arguments_to_service() -> None: + datasets = [ + Dataset( + dataset_uuid="073410b6-79be-3e7d-ae37-92f6e054013e", + study_uuid="64a98a9b-1512-44c8-92af-e4cab0183670", + study_env_uuid="4d1fd10d-5b57-4fd8-a436-f4ec59ce2e4a", + dataset_name="labs", + ) + ] + service = _FakeService(datasets=datasets) + client = DataConnectClient(service) + + result = client.get_datasets( + study_environment_uuid=UUID("4d1fd10d-5b57-4fd8-a436-f4ec59ce2e4a"), + search_dataset_name="labs", + page=2, + page_size=10, + ) + + assert result == datasets + assert service.last_get_datasets_kwargs == { + "study_environment_uuid": UUID("4d1fd10d-5b57-4fd8-a436-f4ec59ce2e4a"), + "search_dataset_name": "labs", + "page": 2, + "page_size": 10, + } + + +def test_get_datasets_uses_defaults() -> None: + service = _FakeService() + client = DataConnectClient(service) + + result = client.get_datasets(study_environment_uuid=UUID("11111111-1111-1111-1111-111111111111")) + + assert result == [] + assert service.last_get_datasets_kwargs == { + "study_environment_uuid": UUID("11111111-1111-1111-1111-111111111111"), + "search_dataset_name": "", + "page": 1, + "page_size": 50, + } diff --git a/tests/test_service.py b/tests/test_service.py index 2873713..65c375a 100644 --- a/tests/test_service.py +++ b/tests/test_service.py @@ -6,7 +6,7 @@ import pytest from dataconnect.exceptions import ConnectionError, ValidationError -from dataconnect.models import DatasetVersion +from dataconnect.models import Dataset, DatasetVersion from dataconnect.service.default import DefaultDataConnectService from dataconnect.transport.errors import TransportConnectionError from dataconnect.transport.models import DataRef, ResourceInfo, ResourceQuery @@ -147,3 +147,87 @@ def test_get_dataset_versions_raises_validation_error_on_zero_input() -> None: # Ensure our validation code path is exercised assert "dataset_uuid must not be empty" in str(excinfo.value) + + +# --- get_datasets tests --- + + +def test_get_datasets_returns_mapped_models_and_builds_request() -> None: + payload = { + "dataset_uuid": "073410b6-79be-3e7d-ae37-92f6e054013e", + "study_uuid": "64a98a9b-1512-44c8-92af-e4cab0183670", + "study_env_uuid": "4d1fd10d-5b57-4fd8-a436-f4ec59ce2e4a", + "dataset_name": "labs", + } + transport = _FakeTransport(resources=[_resource_with_ticket_json(payload)]) + service = DefaultDataConnectService(transport) + + result = service.get_datasets(study_environment_uuid=UUID("4d1fd10d-5b57-4fd8-a436-f4ec59ce2e4a")) + + assert result == [ + Dataset( + dataset_uuid="073410b6-79be-3e7d-ae37-92f6e054013e", + study_uuid="64a98a9b-1512-44c8-92af-e4cab0183670", + study_env_uuid="4d1fd10d-5b57-4fd8-a436-f4ec59ce2e4a", + dataset_name="labs", + ) + ] + assert transport.last_request is not None + assert transport.last_request.action == "datasets.list" + body = json.loads(transport.last_request.body) + assert body["study_environment_uuid"] == "4d1fd10d-5b57-4fd8-a436-f4ec59ce2e4a" + assert body["page"] == 1 + assert body["page_size"] == 50 + + +def test_get_datasets_passes_all_parameters_in_request_body() -> None: + transport = _FakeTransport(resources=[]) + service = DefaultDataConnectService(transport) + + service.get_datasets( + study_environment_uuid=UUID("11111111-1111-1111-1111-111111111111"), + search_dataset_name="vitals", + page=3, + page_size=25, + ) + + assert transport.last_request is not None + body = json.loads(transport.last_request.body) + assert body == { + "study_environment_uuid": "11111111-1111-1111-1111-111111111111", + "search_dataset_name": "vitals", + "page": 3, + "page_size": 25, + } + + +def test_get_datasets_translates_transport_errors() -> None: + transport = _FakeTransport(error=TransportConnectionError("cannot connect")) + service = DefaultDataConnectService(transport) + + with pytest.raises(ConnectionError, match="cannot connect"): + service.get_datasets(study_environment_uuid=UUID("11111111-1111-1111-1111-111111111111")) + + +def test_get_datasets_returns_empty_list_when_no_resources() -> None: + transport = _FakeTransport(resources=[]) + service = DefaultDataConnectService(transport) + + result = service.get_datasets(study_environment_uuid=UUID("11111111-1111-1111-1111-111111111111")) + + assert result == [] + + +def test_get_datasets_raises_validation_error_on_bad_payload() -> None: + # Provide invalid JSON in the ticket to trigger a mapping error + bad_resource = ResourceInfo( + descriptor=b"", + endpoints=[DataRef(ticket=b"not-valid-json")], + total_records=1, + schema_bytes=b"", + ) + transport = _FakeTransport(resources=[bad_resource]) + service = DefaultDataConnectService(transport) + + with pytest.raises(ValidationError, match="Unexpected datasets response format"): + service.get_datasets(study_environment_uuid=UUID("11111111-1111-1111-1111-111111111111"))