diff --git a/dataconnect/service/default.py b/dataconnect/service/default.py index e8f773f..4f81317 100644 --- a/dataconnect/service/default.py +++ b/dataconnect/service/default.py @@ -6,6 +6,7 @@ import pandas as pd +from dataconnect.exceptions import ErrorDetail from dataconnect.models import Dataset, DatasetVersion, PaginatedResponse, Pagination, Study from dataconnect.service.base import DataConnectService from dataconnect.service.error_handler import translate_error @@ -18,13 +19,12 @@ from dataconnect.service.validators import validate_positive_int, validate_uuid from dataconnect.transport.base import Transport from dataconnect.transport.errors import TransportError -from dataconnect.transport.models import ResourceQuery +from dataconnect.transport.models import DatasetTicket, ResourceQuery # Server action identifiers _ACTION_LIST_STUDIES = "studies.list" _ACTION_LIST_DATASETS = "datasets.list" _ACTION_LIST_DATASET_VERSIONS = "dataset_versions.list" -_ACTION_FETCH_TICKET = "data.fetch_ticket" class DefaultDataConnectService(DataConnectService): @@ -85,27 +85,46 @@ def get_dataset_versions(self, dataset_uuid: UUID) -> list[DatasetVersion]: raise translate_error(ex) from ex def fetch_data(self, dataset_uuid: UUID, first_n_rows: int | None = None) -> pd.DataFrame: + """Fetch data for a dataset""" + + validate_uuid( + dataset_uuid, + field_name="dataset_uuid", + error_code="VAL_C_DATASET_UUID", + message="Invalid dataset_uuid.", + details=[ + ErrorDetail( + field="dataset_uuid", + message="dataset_uuid must be a valid UUID.", + expected="Review and provide the correct dataset_uuid.", + ) + ], + ) - if not dataset_uuid or not str(dataset_uuid).strip(): - raise ValueError("dataset_uuid must be provided.") - - if dataset_uuid.int == 0: - raise ValueError("dataset_uuid must not be an empty UUID.") - - if first_n_rows is not None and (not isinstance(first_n_rows, int) or first_n_rows <= 0): - raise ValueError("first_n_rows must be a positive integer when provided.") + if first_n_rows is not None: + validate_positive_int( + first_n_rows, + field_name="first_n_rows", + error_code="VAL_C_FIRST_N_ROWS", + message="Invalid first_n_rows.", + details=[ + ErrorDetail( + field="first_n_rows", + message=(f"Received {first_n_rows} for first_n_rows, which is not a positive integer."), + expected=( + "Set first_n_rows to 1 or greater, or omit the parameter to retrieve the full dataset" + ), + ) + ], + ) - request = ResourceQuery(action=_ACTION_FETCH_TICKET).append_body( - { - "study_env_uuid": None, - "dataset_name": None, - "dataset_uuid": str(dataset_uuid), - "limit": first_n_rows, - } + ticket = DatasetTicket( + dataset_uuid=str(dataset_uuid), + limit=first_n_rows, ) try: - table = self._transport.do_get(request) + table = self._transport.get_ticket(ticket) return resource_to_fetched_data(table) except TransportError as ex: raise translate_error(ex) from ex diff --git a/dataconnect/service/validators.py b/dataconnect/service/validators.py index 9daef24..7fc343a 100644 --- a/dataconnect/service/validators.py +++ b/dataconnect/service/validators.py @@ -5,7 +5,7 @@ from datetime import UTC, datetime from uuid import UUID -from dataconnect.exceptions import ValidationError +from dataconnect.exceptions import ErrorDetail, ValidationError def validate_search_study_name(search_study_name: str | None) -> None: @@ -18,7 +18,14 @@ def validate_search_study_name(search_study_name: str | None) -> None: raise ValidationError("search_study_name must be a string") -def validate_uuid(value: object, *, field_name: str, error_code: str) -> None: +def validate_uuid( + value: object, + *, + field_name: str, + error_code: str, + message: str | None = None, + details: list[ErrorDetail] | None = None, +) -> None: """Ensure *value* is a non-zero UUID. Raises: @@ -27,19 +34,28 @@ def validate_uuid(value: object, *, field_name: str, error_code: str) -> None: if not isinstance(value, UUID): raise ValidationError( error_code=error_code, - message=f"{field_name} must be a valid UUID.", + message=message or f"{field_name} must be a valid UUID.", timestamp=datetime.now(UTC).strftime("%Y-%m-%dT%H:%M:%SZ"), + details=details, ) if value.int == 0: raise ValidationError( error_code=error_code, - message=f"{field_name} must not be empty.", + message=message or f"{field_name} must not be empty.", timestamp=datetime.now(UTC).strftime("%Y-%m-%dT%H:%M:%SZ"), + details=details, ) -def validate_positive_int(value: object, *, field_name: str, error_code: str) -> None: +def validate_positive_int( + value: object, + *, + field_name: str, + error_code: str, + message: str | None = None, + details: list[ErrorDetail] | None = None, +) -> None: """Ensure *value* is an integer >= 1. Raises: @@ -48,6 +64,7 @@ def validate_positive_int(value: object, *, field_name: str, error_code: str) -> if not isinstance(value, int) or value < 1: raise ValidationError( error_code=error_code, - message=f"{field_name} must be a positive integer.", + message=message or f"{field_name} must be a positive integer.", timestamp=datetime.now(UTC).strftime("%Y-%m-%dT%H:%M:%SZ"), + details=details, ) diff --git a/dataconnect/transport/arrow_flight/transport.py b/dataconnect/transport/arrow_flight/transport.py index 27ea45e..801d90b 100644 --- a/dataconnect/transport/arrow_flight/transport.py +++ b/dataconnect/transport/arrow_flight/transport.py @@ -8,6 +8,7 @@ from __future__ import annotations import base64 +import dataclasses import json import platform import subprocess @@ -19,7 +20,7 @@ from dataconnect.transport.arrow_flight.error_handler import parse_dataconnect_error from dataconnect.transport.base import Transport from dataconnect.transport.errors import TransportValidationError -from dataconnect.transport.models import DataRef, DataTable, ResourceInfo, ResourceQuery +from dataconnect.transport.models import DataRef, DatasetTicket, DataTable, ResourceInfo, ResourceQuery def _to_resource_info(info: flight.FlightInfo) -> ResourceInfo: @@ -167,20 +168,10 @@ def list_resources(self, request: ResourceQuery) -> list[ResourceInfo]: except Exception as ex: raise parse_dataconnect_error(ex) from ex - def do_get(self, request: ResourceQuery) -> DataTable: + def get_ticket(self, ticket: DatasetTicket) -> DataTable: """Call FlightClient.do_get and read all chunks into a single pa.Table.""" - flight_type = _ACTION_FLIGHT_TYPE.get(request.action) - - if flight_type is None: - raise TransportValidationError( - error_code="VAL_001", - message=f"Unsupported action: {request.action}", - timestamp=datetime.now(UTC).strftime("%Y-%m-%dT%H:%M:%SZ"), - ) - - body = json.loads(request.body) if request.body else {} - ticket_bytes = json.dumps({**body, "flight_type": flight_type}, separators=(",", ":")).encode("utf-8") + ticket_bytes = json.dumps(dataclasses.asdict(ticket), separators=(",", ":")).encode("utf-8") ticket = flight.Ticket(ticket_bytes) try: diff --git a/dataconnect/transport/base.py b/dataconnect/transport/base.py index 3c06161..939a3fe 100644 --- a/dataconnect/transport/base.py +++ b/dataconnect/transport/base.py @@ -9,7 +9,7 @@ from abc import ABC, abstractmethod -from dataconnect.transport.models import DataTable, ResourceInfo, ResourceQuery +from dataconnect.transport.models import DatasetTicket, DataTable, ResourceInfo, ResourceQuery class Transport(ABC): @@ -24,8 +24,8 @@ def list_resources(self, request: ResourceQuery) -> list[ResourceInfo]: """ @abstractmethod - def do_get(self, request: ResourceQuery) -> DataTable: - """Fetch the full dataset described by ``request``. + def get_ticket(self, ticket: DatasetTicket) -> DataTable: + """Fetch the full dataset described by ``ticket``. All record batches from the stream are read and returned as a single ``DataTable`` containing the complete result set. diff --git a/dataconnect/transport/models.py b/dataconnect/transport/models.py index 6f3027b..6821c36 100644 --- a/dataconnect/transport/models.py +++ b/dataconnect/transport/models.py @@ -30,6 +30,17 @@ class DataRef: ticket: bytes +@dataclass(frozen=True) +class DatasetTicket: + """A data ticket for a specific dataset, containing all information needed to fetch the data.""" + + dataset_uuid: str + limit: int | None = None + study_env_uuid: str | None = None + dataset_name: str | None = None + dataset_version: str | None = None + + @dataclass(frozen=True) class ResourceInfo: """Technology-agnostic representation of a single resource.""" diff --git a/tests/test_fetch_data.py b/tests/test_fetch_data.py new file mode 100644 index 0000000..0f5396e --- /dev/null +++ b/tests/test_fetch_data.py @@ -0,0 +1,242 @@ +"""Unit tests for DefaultDataConnectService.fetch_data.""" + +from __future__ import annotations + +from uuid import UUID + +import pandas as pd +import pyarrow as pa +import pytest + +from dataconnect.exceptions import ( + AuthenticationError, + AuthorizationError, + NotFoundError, + ServerError, + ValidationError, +) +from dataconnect.service.default import DefaultDataConnectService +from dataconnect.transport.errors import ( + TransportAuthenticationError, + TransportAuthorizationError, + TransportNotFoundError, + TransportServerError, +) +from dataconnect.transport.models import DatasetTicket, DataTable, ResourceInfo, ResourceQuery + +# --------------------------------------------------------------------------- +# Fake transport +# --------------------------------------------------------------------------- + + +class _FakeTransport: + """Minimal Transport stub for unit tests.""" + + def __init__( + self, + data_table: DataTable | None = None, + get_ticket_error: Exception | None = None, + ) -> None: + self._data_table = data_table + self._get_ticket_error = get_ticket_error + self.last_ticket: DatasetTicket | None = None + + def list_resources(self, request: ResourceQuery) -> list[ResourceInfo]: + return [] + + def get_ticket(self, ticket: DatasetTicket) -> DataTable: + self.last_ticket = ticket + if self._get_ticket_error is not None: + raise self._get_ticket_error + assert self._data_table is not None + return self._data_table + + def close(self) -> None: + return None + + +def _make_ipc_table(data: dict) -> DataTable: + arrow_table = pa.table(data) + sink = pa.BufferOutputStream() + writer = pa.ipc.new_stream(sink, arrow_table.schema) + writer.write_table(arrow_table) + writer.close() + return DataTable( + schema_bytes=arrow_table.schema.serialize().to_pybytes(), + ipc_bytes=sink.getvalue().to_pybytes(), + ) + + +# --------------------------------------------------------------------------- +# Happy-path +# --------------------------------------------------------------------------- + + +def test_fetch_data_returns_dataframe_with_correct_values() -> None: + dataset_uuid = UUID("073410b6-79be-3e7d-ae37-92f6e054013e") + source = {"subject_id": ["S001", "S002"], "age": [30, 45]} + transport = _FakeTransport(data_table=_make_ipc_table(source)) + service = DefaultDataConnectService(transport) + + result = service.fetch_data(dataset_uuid) + + assert isinstance(result, pd.DataFrame) + assert result["subject_id"].tolist() == source["subject_id"] + assert result["age"].tolist() == source["age"] + + +def test_fetch_data_builds_correct_ticket() -> None: + dataset_uuid = UUID("073410b6-79be-3e7d-ae37-92f6e054013e") + transport = _FakeTransport(data_table=_make_ipc_table({"x": [1]})) + service = DefaultDataConnectService(transport) + + service.fetch_data(dataset_uuid, first_n_rows=10) + + assert transport.last_ticket is not None + assert transport.last_ticket.dataset_uuid == str(dataset_uuid) + assert transport.last_ticket.limit == 10 + + +def test_fetch_data_no_limit_sends_none_in_ticket() -> None: + dataset_uuid = UUID("073410b6-79be-3e7d-ae37-92f6e054013e") + transport = _FakeTransport(data_table=_make_ipc_table({"x": [1]})) + service = DefaultDataConnectService(transport) + + service.fetch_data(dataset_uuid) + + assert transport.last_ticket is not None + assert transport.last_ticket.limit is None + + +def test_fetch_data_returns_empty_dataframe_for_empty_table() -> None: + dataset_uuid = UUID("073410b6-79be-3e7d-ae37-92f6e054013e") + transport = _FakeTransport(data_table=_make_ipc_table({"col": pa.array([], type=pa.int32())})) + service = DefaultDataConnectService(transport) + + result = service.fetch_data(dataset_uuid) + + assert isinstance(result, pd.DataFrame) + assert len(result) == 0 + assert "col" in result.columns + + +# --------------------------------------------------------------------------- +# Validation — dataset_uuid +# --------------------------------------------------------------------------- + + +def test_fetch_data_raises_validation_error_on_none_uuid() -> None: + transport = _FakeTransport(data_table=_make_ipc_table({"x": [1]})) + service = DefaultDataConnectService(transport) + + with pytest.raises(ValidationError) as exc_info: + service.fetch_data(None) # type: ignore[arg-type] + + assert exc_info.value.error_code == "VAL_C_DATASET_UUID" + + +def test_fetch_data_raises_validation_error_on_zero_uuid() -> None: + transport = _FakeTransport(data_table=_make_ipc_table({"x": [1]})) + service = DefaultDataConnectService(transport) + + with pytest.raises(ValidationError) as exc_info: + service.fetch_data(UUID(int=0)) + + assert exc_info.value.error_code == "VAL_C_DATASET_UUID" + + +def test_fetch_data_raises_validation_error_on_string_uuid() -> None: + transport = _FakeTransport(data_table=_make_ipc_table({"x": [1]})) + service = DefaultDataConnectService(transport) + + with pytest.raises(ValidationError) as exc_info: + service.fetch_data("073410b6-79be-3e7d-ae37-92f6e054013e") # type: ignore[arg-type] + + assert exc_info.value.error_code == "VAL_C_DATASET_UUID" + + +# --------------------------------------------------------------------------- +# Validation — first_n_rows +# --------------------------------------------------------------------------- + + +def test_fetch_data_raises_validation_error_on_zero_first_n_rows() -> None: + dataset_uuid = UUID("073410b6-79be-3e7d-ae37-92f6e054013e") + transport = _FakeTransport(data_table=_make_ipc_table({"x": [1]})) + service = DefaultDataConnectService(transport) + + with pytest.raises(ValidationError) as exc_info: + service.fetch_data(dataset_uuid, first_n_rows=0) + + assert exc_info.value.error_code == "VAL_C_FIRST_N_ROWS" + assert exc_info.value.details is not None + assert exc_info.value.details[0].field == "first_n_rows" + + +def test_fetch_data_raises_validation_error_on_negative_first_n_rows() -> None: + dataset_uuid = UUID("073410b6-79be-3e7d-ae37-92f6e054013e") + transport = _FakeTransport(data_table=_make_ipc_table({"x": [1]})) + service = DefaultDataConnectService(transport) + + with pytest.raises(ValidationError) as exc_info: + service.fetch_data(dataset_uuid, first_n_rows=-5) + + assert exc_info.value.error_code == "VAL_C_FIRST_N_ROWS" + + +def test_fetch_data_raises_validation_error_on_non_int_first_n_rows() -> None: + dataset_uuid = UUID("073410b6-79be-3e7d-ae37-92f6e054013e") + transport = _FakeTransport(data_table=_make_ipc_table({"x": [1]})) + service = DefaultDataConnectService(transport) + + with pytest.raises(ValidationError) as exc_info: + service.fetch_data(dataset_uuid, first_n_rows="abc") # type: ignore[arg-type] + + assert exc_info.value.error_code == "VAL_C_FIRST_N_ROWS" + + +# --------------------------------------------------------------------------- +# Transport-error translation +# --------------------------------------------------------------------------- + + +def test_fetch_data_translates_authentication_error() -> None: + dataset_uuid = UUID("073410b6-79be-3e7d-ae37-92f6e054013e") + transport = _FakeTransport( + get_ticket_error=TransportAuthenticationError(error_code="AUTH_E_001", message="bad token") + ) + service = DefaultDataConnectService(transport) + + with pytest.raises(AuthenticationError): + service.fetch_data(dataset_uuid) + + +def test_fetch_data_translates_authorization_error() -> None: + dataset_uuid = UUID("073410b6-79be-3e7d-ae37-92f6e054013e") + transport = _FakeTransport( + get_ticket_error=TransportAuthorizationError(error_code="AUTHZ_001", message="forbidden") + ) + service = DefaultDataConnectService(transport) + + with pytest.raises(AuthorizationError): + service.fetch_data(dataset_uuid) + + +def test_fetch_data_translates_not_found_error() -> None: + dataset_uuid = UUID("073410b6-79be-3e7d-ae37-92f6e054013e") + transport = _FakeTransport( + get_ticket_error=TransportNotFoundError(error_code="RES_001", message="dataset not found") + ) + service = DefaultDataConnectService(transport) + + with pytest.raises(NotFoundError): + service.fetch_data(dataset_uuid) + + +def test_fetch_data_translates_server_error() -> None: + dataset_uuid = UUID("073410b6-79be-3e7d-ae37-92f6e054013e") + transport = _FakeTransport(get_ticket_error=TransportServerError(error_code="INT_001", message="internal error")) + service = DefaultDataConnectService(transport) + + with pytest.raises(ServerError): + service.fetch_data(dataset_uuid)