diff --git a/dataconnect/__init__.py b/dataconnect/__init__.py index 6f9f0a0..e676ca1 100644 --- a/dataconnect/__init__.py +++ b/dataconnect/__init__.py @@ -11,7 +11,7 @@ ServerError, ValidationError, ) -from dataconnect.models import DatasetVersion, PaginatedResponse, Pagination, Study, StudyEnvironment +from dataconnect.models import DatasetVersion, DryPublishResult, PaginatedResponse, Pagination, Study, StudyEnvironment __all__ = [ # Client @@ -20,6 +20,7 @@ "Study", "StudyEnvironment", "DatasetVersion", + "DryPublishResult", "PaginatedResponse", "Pagination", # Exceptions — catch these in user application code diff --git a/dataconnect/client.py b/dataconnect/client.py index c510b44..2a59da5 100644 --- a/dataconnect/client.py +++ b/dataconnect/client.py @@ -12,7 +12,7 @@ import pandas as pd -from dataconnect.models import Dataset, DatasetVersion, PaginatedResponse, StudiesResult +from dataconnect.models import Dataset, DatasetVersion, DryPublishResult, PaginatedResponse, StudiesResult from dataconnect.service import DataConnectService, DefaultDataConnectService _DEFAULT_HOST = "enodia-gateway.platform.imedidata.com" @@ -87,6 +87,25 @@ def get_datasets( page_size=page_size, ) + def dry_publish( + self, + project_token: str, + dataset_name: str, + key_columns: list[str], + source_datasets: list[UUID], + data: pd.DataFrame, + datetime_formats: dict[str, str] | None = None, + ) -> DryPublishResult: + + return self._service.dry_publish( + project_token=project_token, + dataset_name=dataset_name, + key_columns=key_columns, + source_datasets=source_datasets, + data=data, + datetime_formats=datetime_formats, + ) + # Lifecycle def close(self) -> None: diff --git a/dataconnect/models.py b/dataconnect/models.py index 53bc856..8677946 100644 --- a/dataconnect/models.py +++ b/dataconnect/models.py @@ -4,6 +4,8 @@ from typing import Generic, TypeVar from uuid import UUID +import pandas as pd + T = TypeVar("T") @@ -61,3 +63,22 @@ class PaginatedResponse(Generic[T]): # noqa: UP046 total_records: int pagination: Pagination items: list[T] + + +@dataclass +class DryPublishResult: + """Result of a dry publish operation, including validation status and details.""" + + status: bool + is_schema_valid: bool | None = None + is_config_valid: bool | None = None + is_dataset_valid: bool | None = None + errors: list[str] = field(default_factory=list) + invalid_datetime_formats: dict[str, str] = field(default_factory=dict) + dataset_name: str | None = None + dataset_version: int | None = None + no_of_columns: int | None = None + valid_record_count: int | None = None + duplicate_record_count: int | None = None + invalid_record_count: int | None = None + invalid_records: pd.DataFrame | None = None diff --git a/dataconnect/service/base.py b/dataconnect/service/base.py index c85fae6..c315cd4 100644 --- a/dataconnect/service/base.py +++ b/dataconnect/service/base.py @@ -7,7 +7,7 @@ import pandas as pd -from dataconnect.models import Dataset, DatasetVersion, PaginatedResponse, StudiesResult +from dataconnect.models import Dataset, DatasetVersion, DryPublishResult, PaginatedResponse, StudiesResult class DataConnectService(ABC): @@ -35,5 +35,16 @@ def fetch_data( first_n_rows: int | None = None, ) -> pd.DataFrame: ... + @abstractmethod + def dry_publish( + self, + project_token: str, + dataset_name: str, + key_columns: list[str], + source_datasets: list[UUID], + data: pd.DataFrame, + datetime_formats: dict[str, str] | None = None, + ) -> DryPublishResult: ... + @abstractmethod def close(self) -> None: ... diff --git a/dataconnect/service/default.py b/dataconnect/service/default.py index 72dfa74..0e96f55 100644 --- a/dataconnect/service/default.py +++ b/dataconnect/service/default.py @@ -2,14 +2,16 @@ from __future__ import annotations +import json from uuid import UUID import pandas as pd -from dataconnect.models import Dataset, DatasetVersion, PaginatedResponse, Pagination, StudiesResult +from dataconnect.models import Dataset, DatasetVersion, DryPublishResult, PaginatedResponse, Pagination, StudiesResult from dataconnect.service.base import DataConnectService from dataconnect.service.error_handler import translate_error from dataconnect.service.mappers import ( + dry_publish_response_to_domain, resource_to_dataset, resource_to_dataset_version, resource_to_fetched_data, @@ -17,7 +19,7 @@ ) from dataconnect.transport.base import Transport from dataconnect.transport.errors import TransportError -from dataconnect.transport.models import DatasetTicket, ResourceQuery +from dataconnect.transport.models import DatasetTicket, PublishRequest, ResourceQuery # Server action identifiers _ACTION_LIST_STUDIES = "studies.list" @@ -138,6 +140,71 @@ def get_datasets( except TransportError as ex: raise translate_error(ex) from ex + def dry_publish( + self, + project_token: str, + dataset_name: str, + key_columns: list[str], + source_datasets: list[UUID], + data: pd.DataFrame, + datetime_formats: dict[str, str] | None = None, + ) -> DryPublishResult: + """Validate a dataset against the server without committing any changes. + + Encodes the publish configuration as a JSON ``input_config`` string, + sets ``is_dry_publish: True`` so the server treats the call as a + validation-only run, then delegates to the transport layer. + + The server response is mapped to a :class:`DryPublishResult` via + :func:`dry_publish_response_to_domain`. + + Args: + project_token: Base64-encoded project token identifying the target + study, study environment, and project. + dataset_name: Name of the dataset to validate. + key_columns: Column names that form the unique key for the dataset. + source_datasets: UUIDs of the source datasets the published dataset + is derived from. + data: The dataset to validate as a ``pd.DataFrame``. + datetime_formats: Optional mapping of column name → datetime format + string (e.g. ``{"visit_date": "yyyy-MM-dd"}``). Defaults to + an empty dict when omitted. + + Returns: + A :class:`DryPublishResult` containing the server's validation + outcome, including per-field validity flags, error messages, and + an optional ``invalid_records`` DataFrame. + + Raises: + DataConnectError: Any :class:`TransportError` from the transport + layer is translated by :func:`translate_error` into the public + API's :class:`DataConnectError` hierarchy before propagating, + for example :class:`ValidationError`, + :class:`AuthorizationError`, or :class:`ServerError`.""" + + request = PublishRequest( + input_config=json.dumps( + { + "is_dry_publish": True, + "project_token": project_token, + "dataset_name": dataset_name, + "dataset_description": dataset_name, + "key_columns": key_columns, + "source_datasets": [str(uuid) for uuid in source_datasets], + "datetime_formats": datetime_formats or {}, + } + ), + data=data, + ) + + try: + publish_result = self._transport.dry_publish_dataset(request) + + return dry_publish_response_to_domain(publish_result) + + except TransportError as ex: + raise translate_error(ex) from ex + def close(self) -> None: """Close the underlying transport connection.""" diff --git a/dataconnect/service/mappers.py b/dataconnect/service/mappers.py index 8141fbb..0a8513e 100644 --- a/dataconnect/service/mappers.py +++ b/dataconnect/service/mappers.py @@ -15,8 +15,8 @@ import pyarrow as pa from dataconnect.exceptions import NotFoundError -from dataconnect.models import Dataset, DatasetVersion, Study, StudyEnvironment -from dataconnect.transport.models import DataTable, ResourceInfo +from dataconnect.models import Dataset, DatasetVersion, DryPublishResult, Study, StudyEnvironment +from dataconnect.transport.models import DataTable, DryPublishResponse, ResourceInfo def resource_to_study(resource: ResourceInfo) -> Study: @@ -85,3 +85,42 @@ def resource_to_dataset(resource: ResourceInfo) -> Dataset: study_env_uuid=data.get("study_env_uuid", ""), dataset_name=data.get("dataset_name", ""), ) + + +def dry_publish_response_to_domain(result: DryPublishResponse | None) -> DryPublishResult: + """Map a transport-layer ``DryPublishResponse`` to a ``DryPublishResult`` domain object. + + ``DryPublishResponse`` carries flat, typed fields returned by the server after a + dry-publish call. The mapping is direct for all shared fields with one + exception: + + * ``DryPublishResponse.dataset_valid`` → ``DryPublishResult.is_dataset_valid`` + (renamed for naming consistency with the other ``is_*_valid`` fields). + + Args: + result: The transport-layer result returned by + :meth:`Transport.dry_publish_dataset`. Pass ``None`` to obtain a + default :class:`DryPublishResult` with ``status=False`` and all + other fields at their zero values. + + Returns: + A :class:`DryPublishResult` suitable for returning to the caller. + """ + if result is None: + return DryPublishResult(status=False) + + return DryPublishResult( + status=result.status, + is_schema_valid=result.is_schema_valid, + is_config_valid=result.is_config_valid, + is_dataset_valid=result.dataset_valid, + errors=result.errors, + invalid_datetime_formats=result.invalid_datetime_formats, + dataset_name=result.dataset_name, + dataset_version=result.dataset_version, + no_of_columns=result.no_of_columns, + valid_record_count=result.valid_record_count, + duplicate_record_count=result.duplicate_record_count, + invalid_record_count=result.invalid_record_count, + invalid_records=result.invalid_records, + ) diff --git a/dataconnect/transport/__init__.py b/dataconnect/transport/__init__.py index 3402521..d3024f2 100644 --- a/dataconnect/transport/__init__.py +++ b/dataconnect/transport/__init__.py @@ -6,11 +6,23 @@ """ from dataconnect.transport.base import Transport -from dataconnect.transport.models import DataRef, ResourceInfo, ResourceQuery +from dataconnect.transport.models import ( + DataRef, + DatasetTicket, + DataTable, + DryPublishResponse, + PublishRequest, + ResourceInfo, + ResourceQuery, +) __all__ = [ "Transport", "ResourceQuery", "ResourceInfo", "DataRef", + "DatasetTicket", + "DataTable", + "PublishRequest", + "DryPublishResponse", ] diff --git a/dataconnect/transport/arrow_flight/transport.py b/dataconnect/transport/arrow_flight/transport.py index 4458010..d33b562 100644 --- a/dataconnect/transport/arrow_flight/transport.py +++ b/dataconnect/transport/arrow_flight/transport.py @@ -22,7 +22,15 @@ 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, DatasetTicket, DataTable, ResourceInfo, ResourceQuery +from dataconnect.transport.models import ( + DataRef, + DatasetTicket, + DataTable, + DryPublishResponse, + PublishRequest, + ResourceInfo, + ResourceQuery, +) def _to_resource_info(info: flight.FlightInfo) -> ResourceInfo: @@ -59,6 +67,25 @@ def _to_bytes(table: pa.Table) -> DataTable: return DataTable(schema_bytes=schema_bytes, ipc_bytes=ipc_bytes) +def _normalize_arrow_type(dtype: pa.DataType) -> pa.DataType: + """Recursively normalize Arrow types that widen during a pandas round-trip. + + ``pa.Table.from_pandas()`` always infers the "large" variants because pandas + has no distinction between them: + + * ``large_string`` → ``string`` + * ``large_binary`` → ``binary`` + * ``large_list`` → ``list`` (applied recursively to the value type) + """ + if dtype == pa.large_utf8(): + return pa.utf8() + if dtype == pa.large_binary(): + return pa.binary() + if pa.types.is_large_list(dtype): + return pa.list_(_normalize_arrow_type(dtype.value_type)) + return dtype + + # Maps service-layer action names to the flight_type value the Arrow Flight server expects. _ACTION_FLIGHT_TYPE: dict[str, str] = { "studies.list": "STUDIES", @@ -214,5 +241,85 @@ def get_ticket(self, ticket: DatasetTicket) -> DataTable: except Exception as ex: raise parse_dataconnect_error(ex) from ex + def dry_publish_dataset(self, publish_request: PublishRequest) -> DryPublishResponse: + """Send a dataset to the server via Arrow Flight ``do_put`` and return the validation result. + + The method serialises the request DataFrame to Arrow IPC, streams it to + the server batch-by-batch, then reads two metadata responses from the + server before closing the call: + + 1. A JSON buffer containing validation fields (status, error lists, …). + 2. An optional Arrow IPC buffer containing the invalid-records table + (``None`` when all rows pass validation). + + The ``done_writing()`` / ``close()`` split is intentional: + * ``done_writing()`` signals end-of-stream without closing the RPC call, + allowing the metadata reads to complete while the call is still alive. + * ``writer.close()`` in the ``finally`` block always terminates the call, + even if writing or reading raises. + + Args: + request: A :class:`PublishRequest` carrying the encoded server config + (``input_config``) and the dataset as a ``pd.DataFrame``. + + Returns: + A :class:`DryPublishResponse` populated from the server's JSON response and, + when present, the invalid-records Arrow table converted to a + ``pd.DataFrame``. + + Raises: + TransportError: Any Arrow Flight or gRPC error is translated by + :func:`parse_dataconnect_error` before propagating. + """ + + descriptor_bytes = publish_request.input_config.encode("utf-8") + flight_descriptor = flight.FlightDescriptor.for_path(descriptor_bytes) + + arrow_table = pa.Table.from_pandas(publish_request.data, preserve_index=False) + + # pa.Table.from_pandas() widens string→large_string, binary→large_binary, + # and list→large_list. Normalize back so the schema matches the server. + arrow_table = arrow_table.cast( + pa.schema([f.with_type(_normalize_arrow_type(f.type)) for f in arrow_table.schema]) + ) + + try: + writer, reader = self._client.do_put(flight_descriptor, arrow_table.schema, self._options()) + + try: + for batch in arrow_table.to_batches(): + writer.write_batch(batch) + writer.done_writing() # only signal completion when all batches succeeded + + # Read while the RPC call is still open (before writer.close()) + # The server first writes a JSON result, then the Arrow table as IPC bytes + _json_buf = reader.read() + json_result = json.loads(_json_buf.to_pybytes()) + + # Read the Arrow table returned by the server upon successful publishing + metadata_buf = reader.read() + result_table = pa.ipc.open_stream(pa.BufferReader(metadata_buf)).read_all() if metadata_buf else None + finally: + writer.close() # terminates the RPC call — must happen after all reads + + return DryPublishResponse( + status=json_result.get("status", False), + is_schema_valid=json_result.get("is_schema_valid", False), + is_config_valid=json_result.get("is_config_valid", False), + dataset_valid=json_result.get("dataset_valid", False), + errors=json_result.get("errors", []), + invalid_datetime_formats=json_result.get("invalid_datetime_formats", {}), + dataset_name=json_result.get("dataset_name", ""), + dataset_version=json_result.get("dataset_version", 0), + no_of_columns=json_result.get("no_of_columns", 0), + valid_record_count=json_result.get("valid_record_count", 0), + duplicate_record_count=json_result.get("duplicate_record_count", 0), + invalid_record_count=json_result.get("invalid_record_count", 0), + invalid_records=result_table.to_pandas() if result_table else None, + ) + + except Exception as ex: + raise parse_dataconnect_error(ex) from ex + def close(self) -> None: self._client.close() diff --git a/dataconnect/transport/base.py b/dataconnect/transport/base.py index 939a3fe..da668ed 100644 --- a/dataconnect/transport/base.py +++ b/dataconnect/transport/base.py @@ -9,7 +9,14 @@ from abc import ABC, abstractmethod -from dataconnect.transport.models import DatasetTicket, DataTable, ResourceInfo, ResourceQuery +from dataconnect.transport.models import ( + DatasetTicket, + DataTable, + DryPublishResponse, + PublishRequest, + ResourceInfo, + ResourceQuery, +) class Transport(ABC): @@ -31,6 +38,14 @@ def get_ticket(self, ticket: DatasetTicket) -> DataTable: ``DataTable`` containing the complete result set. """ + @abstractmethod + def dry_publish_dataset(self, publish_request: PublishRequest) -> DryPublishResponse: + """Dry publish a dataset described by ``publish_request``. + + All record batches from the stream are read and returned as a single + ``DryPublishResponse`` containing the complete result set. + """ + @abstractmethod def close(self) -> None: """Close the transport connection.""" diff --git a/dataconnect/transport/models.py b/dataconnect/transport/models.py index 6821c36..348e7b6 100644 --- a/dataconnect/transport/models.py +++ b/dataconnect/transport/models.py @@ -6,6 +6,8 @@ from dataclasses import dataclass, field from typing import Any +import pandas as pd + @dataclass(frozen=True) class ResourceQuery: @@ -61,3 +63,26 @@ class DataTable: schema_bytes: bytes ipc_bytes: bytes + + +@dataclass(frozen=True) +class PublishRequest: + input_config: str + data: pd.DataFrame + + +@dataclass(frozen=True) +class DryPublishResponse: + status: bool + is_schema_valid: bool + is_config_valid: bool + dataset_valid: bool + errors: list[str] + invalid_datetime_formats: dict[str, str] + dataset_name: str + dataset_version: int + no_of_columns: int + valid_record_count: int + duplicate_record_count: int + invalid_record_count: int = 0 + invalid_records: pd.DataFrame | None = None diff --git a/tests/test_dry_publish.py b/tests/test_dry_publish.py new file mode 100644 index 0000000..316525b --- /dev/null +++ b/tests/test_dry_publish.py @@ -0,0 +1,513 @@ +"""Unit tests for the dry-publish feature. + +Covers: +- ``_normalize_arrow_type`` (transport helper) +- ``dry_publish_response_to_domain`` (service mapper) +- ``DefaultDataConnectService.dry_publish`` (service layer) +- ``ArrowFlightTransport.dry_publish_dataset`` (transport layer) +""" + +from __future__ import annotations + +import json +from unittest.mock import MagicMock, patch +from uuid import UUID + +import pandas as pd +import pyarrow as pa +import pytest + +from dataconnect.exceptions import ValidationError +from dataconnect.models import DryPublishResult +from dataconnect.service.default import DefaultDataConnectService +from dataconnect.service.mappers import dry_publish_response_to_domain +from dataconnect.transport.arrow_flight.transport import ( + ArrowFlightTransport, + _normalize_arrow_type, +) +from dataconnect.transport.base import Transport +from dataconnect.transport.errors import TransportValidationError +from dataconnect.transport.models import ( + DatasetTicket, + DataTable, + DryPublishResponse, + PublishRequest, + ResourceInfo, + ResourceQuery, +) + +# --------------------------------------------------------------------------- +# Helpers shared across test suites +# --------------------------------------------------------------------------- + + +def _make_dry_publish_response(**overrides: object) -> DryPublishResponse: + """Return a fully-populated ``DryPublishResponse`` with sensible defaults.""" + defaults: dict = dict( + status=True, + is_schema_valid=True, + is_config_valid=True, + dataset_valid=True, + errors=[], + invalid_datetime_formats={}, + dataset_name="demo_dataset", + dataset_version=1, + no_of_columns=5, + valid_record_count=10, + duplicate_record_count=0, + invalid_record_count=0, + invalid_records=None, + ) + return DryPublishResponse(**{**defaults, **overrides}) + + +def _make_json_buf(d: dict) -> pa.Buffer: + return pa.py_buffer(json.dumps(d).encode("utf-8")) + + +def _make_ipc_buf(df: pd.DataFrame) -> pa.Buffer: + table = pa.Table.from_pandas(df) + sink = pa.BufferOutputStream() + w = pa.ipc.new_stream(sink, table.schema) + w.write_table(table) + w.close() + return pa.py_buffer(sink.getvalue().to_pybytes()) + + +# --------------------------------------------------------------------------- +# _normalize_arrow_type +# --------------------------------------------------------------------------- + + +class TestNormalizeArrowType: + """``_normalize_arrow_type`` must downcast the three widened types and leave + everything else unchanged.""" + + def test_large_utf8_becomes_utf8(self) -> None: + assert _normalize_arrow_type(pa.large_utf8()) == pa.utf8() + + def test_large_binary_becomes_binary(self) -> None: + assert _normalize_arrow_type(pa.large_binary()) == pa.binary() + + def test_large_list_of_large_string_becomes_list_of_string(self) -> None: + result = _normalize_arrow_type(pa.large_list(pa.large_utf8())) + assert result == pa.list_(pa.utf8()) + + def test_large_list_recursive_normalization_with_large_binary(self) -> None: + result = _normalize_arrow_type(pa.large_list(pa.large_binary())) + assert result == pa.list_(pa.binary()) + + def test_regular_utf8_passes_through(self) -> None: + assert _normalize_arrow_type(pa.utf8()) == pa.utf8() + + def test_regular_binary_passes_through(self) -> None: + assert _normalize_arrow_type(pa.binary()) == pa.binary() + + def test_regular_list_passes_through(self) -> None: + assert _normalize_arrow_type(pa.list_(pa.utf8())) == pa.list_(pa.utf8()) + + def test_int64_passes_through(self) -> None: + assert _normalize_arrow_type(pa.int64()) == pa.int64() + + def test_float64_passes_through(self) -> None: + assert _normalize_arrow_type(pa.float64()) == pa.float64() + + def test_timestamp_passes_through(self) -> None: + assert _normalize_arrow_type(pa.timestamp("us")) == pa.timestamp("us") + + def test_bool_passes_through(self) -> None: + assert _normalize_arrow_type(pa.bool_()) == pa.bool_() + + +# --------------------------------------------------------------------------- +# dry_publish_response_to_domain +# --------------------------------------------------------------------------- + + +class TestDryPublishResponseToDomain: + """Every field of ``DryPublishResponse`` must land on the correct attribute + of ``DryPublishResult``, including the ``dataset_valid`` rename.""" + + def test_status_true_is_mapped(self) -> None: + assert dry_publish_response_to_domain(_make_dry_publish_response(status=True)).status is True + + def test_status_false_is_mapped(self) -> None: + assert dry_publish_response_to_domain(_make_dry_publish_response(status=False)).status is False + + def test_is_schema_valid_mapped(self) -> None: + result = dry_publish_response_to_domain(_make_dry_publish_response(is_schema_valid=False)) + assert result.is_schema_valid is False + + def test_is_config_valid_mapped(self) -> None: + result = dry_publish_response_to_domain(_make_dry_publish_response(is_config_valid=False)) + assert result.is_config_valid is False + + def test_dataset_valid_renamed_to_is_dataset_valid(self) -> None: + # Transport field name: dataset_valid + # Domain field name: is_dataset_valid + result = dry_publish_response_to_domain(_make_dry_publish_response(dataset_valid=False)) + assert result.is_dataset_valid is False + + def test_errors_list_mapped(self) -> None: + errors = ["Missing column X", "Type mismatch on Y"] + result = dry_publish_response_to_domain(_make_dry_publish_response(errors=errors)) + assert result.errors == errors + + def test_empty_errors_list_mapped(self) -> None: + result = dry_publish_response_to_domain(_make_dry_publish_response(errors=[])) + assert result.errors == [] + + def test_invalid_datetime_formats_mapped(self) -> None: + fmt = {"visit_date": "yyyy-MM-dd"} + result = dry_publish_response_to_domain(_make_dry_publish_response(invalid_datetime_formats=fmt)) + assert result.invalid_datetime_formats == fmt + + def test_dataset_name_mapped(self) -> None: + result = dry_publish_response_to_domain(_make_dry_publish_response(dataset_name="my_ds")) + assert result.dataset_name == "my_ds" + + def test_dataset_version_mapped(self) -> None: + result = dry_publish_response_to_domain(_make_dry_publish_response(dataset_version=42)) + assert result.dataset_version == 42 + + def test_no_of_columns_mapped(self) -> None: + result = dry_publish_response_to_domain(_make_dry_publish_response(no_of_columns=7)) + assert result.no_of_columns == 7 + + def test_valid_record_count_mapped(self) -> None: + result = dry_publish_response_to_domain(_make_dry_publish_response(valid_record_count=100)) + assert result.valid_record_count == 100 + + def test_duplicate_record_count_mapped(self) -> None: + result = dry_publish_response_to_domain(_make_dry_publish_response(duplicate_record_count=3)) + assert result.duplicate_record_count == 3 + + def test_invalid_record_count_mapped(self) -> None: + result = dry_publish_response_to_domain(_make_dry_publish_response(invalid_record_count=2)) + assert result.invalid_record_count == 2 + + def test_invalid_records_none_preserved(self) -> None: + result = dry_publish_response_to_domain(_make_dry_publish_response(invalid_records=None)) + assert result.invalid_records is None + + def test_invalid_records_dataframe_preserved(self) -> None: + df = pd.DataFrame({"col": [1, 2, 3]}) + result = dry_publish_response_to_domain(_make_dry_publish_response(invalid_records=df)) + assert result.invalid_records is df + + def test_returns_dry_publish_result_instance(self) -> None: + result = dry_publish_response_to_domain(_make_dry_publish_response()) + assert isinstance(result, DryPublishResult) + + +# --------------------------------------------------------------------------- +# DefaultDataConnectService.dry_publish +# --------------------------------------------------------------------------- + + +class _StubTransport(Transport): + """Minimal stub that satisfies the Transport ABC for dry-publish tests.""" + + def __init__( + self, + dry_publish_return: DryPublishResponse | None = None, + raise_error: Exception | None = None, + ) -> None: + self._return = dry_publish_return + self._raise = raise_error + self.last_request: PublishRequest | None = None + + def list_resources(self, request: ResourceQuery) -> list[ResourceInfo]: + return [] + + def get_ticket(self, ticket: DatasetTicket) -> DataTable: + raise NotImplementedError + + def dry_publish_dataset(self, request: PublishRequest) -> DryPublishResponse: + self.last_request = request + if self._raise is not None: + raise self._raise + return self._return # type: ignore[return-value] + + def close(self) -> None: + pass + + +def _make_service( + dry_publish_return: DryPublishResponse | None = None, + raise_error: Exception | None = None, +) -> tuple[DefaultDataConnectService, _StubTransport]: + transport = _StubTransport(dry_publish_return=dry_publish_return, raise_error=raise_error) + return DefaultDataConnectService(transport), transport + + +def _default_dry_publish_args() -> dict: + return dict( + project_token="tok_abc123", + dataset_name="my_dataset", + key_columns=["site_id"], + source_datasets=[UUID("0158ea12-4004-3817-899b-2de6becbc0f9")], + data=pd.DataFrame({"site_id": ["S1"], "value": [1]}), + ) + + +class TestDryPublishService: + """``DefaultDataConnectService.dry_publish`` must build the correct request, + delegate to the transport, map the response, and translate errors.""" + + def test_returns_dry_publish_result_instance(self) -> None: + service, _ = _make_service(dry_publish_return=_make_dry_publish_response()) + result = service.dry_publish(**_default_dry_publish_args()) + assert isinstance(result, DryPublishResult) + + def test_successful_status_is_propagated(self) -> None: + service, _ = _make_service(dry_publish_return=_make_dry_publish_response(status=True)) + assert service.dry_publish(**_default_dry_publish_args()).status is True + + # --- input_config JSON content --- + + def test_input_config_sets_is_dry_publish_true(self) -> None: + service, transport = _make_service(dry_publish_return=_make_dry_publish_response()) + service.dry_publish(**_default_dry_publish_args()) + assert transport.last_request is not None + config = json.loads(transport.last_request.input_config) + assert config["is_dry_publish"] is True + + def test_input_config_contains_project_token(self) -> None: + service, transport = _make_service(dry_publish_return=_make_dry_publish_response()) + service.dry_publish(**_default_dry_publish_args()) + assert transport.last_request is not None + assert json.loads(transport.last_request.input_config)["project_token"] == "tok_abc123" + + def test_input_config_contains_dataset_name(self) -> None: + service, transport = _make_service(dry_publish_return=_make_dry_publish_response()) + service.dry_publish(**_default_dry_publish_args()) + assert transport.last_request is not None + assert json.loads(transport.last_request.input_config)["dataset_name"] == "my_dataset" + + def test_input_config_contains_key_columns(self) -> None: + service, transport = _make_service(dry_publish_return=_make_dry_publish_response()) + service.dry_publish(**_default_dry_publish_args()) + assert transport.last_request is not None + assert json.loads(transport.last_request.input_config)["key_columns"] == ["site_id"] + + def test_input_config_source_datasets_serialized_as_strings(self) -> None: + service, transport = _make_service(dry_publish_return=_make_dry_publish_response()) + service.dry_publish(**_default_dry_publish_args()) + assert transport.last_request is not None + config = json.loads(transport.last_request.input_config) + assert config["source_datasets"] == ["0158ea12-4004-3817-899b-2de6becbc0f9"] + + def test_datetime_formats_defaults_to_empty_dict(self) -> None: + service, transport = _make_service(dry_publish_return=_make_dry_publish_response()) + service.dry_publish(**_default_dry_publish_args()) # no datetime_formats kwarg + assert transport.last_request is not None + assert json.loads(transport.last_request.input_config)["datetime_formats"] == {} + + def test_datetime_formats_passed_through_when_provided(self) -> None: + service, transport = _make_service(dry_publish_return=_make_dry_publish_response()) + args = _default_dry_publish_args() + args["datetime_formats"] = {"visit_date": "yyyy-MM-dd"} + service.dry_publish(**args) + assert transport.last_request is not None + assert json.loads(transport.last_request.input_config)["datetime_formats"] == {"visit_date": "yyyy-MM-dd"} + + def test_data_dataframe_passed_to_transport(self) -> None: + service, transport = _make_service(dry_publish_return=_make_dry_publish_response()) + df = pd.DataFrame({"a": [1, 2, 3]}) + args = _default_dry_publish_args() + args["data"] = df + service.dry_publish(**args) + assert transport.last_request is not None + assert transport.last_request.data is df + + # --- falsy / None result --- + + def test_none_transport_result_returns_status_false(self) -> None: + service, _ = _make_service(dry_publish_return=None) + result = service.dry_publish(**_default_dry_publish_args()) + assert isinstance(result, DryPublishResult) + assert result.status is False + + # --- error translation --- + + def test_transport_validation_error_is_translated_to_service_error(self) -> None: + err = TransportValidationError( + error_code="VAL_001", + message="schema mismatch", + timestamp="2024-01-01T00:00:00Z", + ) + service, _ = _make_service(raise_error=err) + with pytest.raises(ValidationError): + service.dry_publish(**_default_dry_publish_args()) + + +# --------------------------------------------------------------------------- +# ArrowFlightTransport.dry_publish_dataset +# --------------------------------------------------------------------------- + + +def _make_flight_transport() -> ArrowFlightTransport: + """Create an ``ArrowFlightTransport`` with a mocked ``FlightClient``.""" + with patch.object(ArrowFlightTransport, "_get_client", return_value=MagicMock()): + return ArrowFlightTransport(host="localhost", port=5005, use_tls=False) + + +def _wire_do_put( + transport: ArrowFlightTransport, + json_resp: dict, + invalid_records_df: pd.DataFrame | None = None, +) -> tuple[MagicMock, MagicMock]: + """Configure ``transport._client.do_put`` to return controlled writer/reader mocks.""" + json_buf = _make_json_buf(json_resp) + ipc_buf = _make_ipc_buf(invalid_records_df) if invalid_records_df is not None else None + + reader_mock = MagicMock() + reader_mock.read.side_effect = [json_buf, ipc_buf] + writer_mock = MagicMock() + transport._client.do_put.return_value = (writer_mock, reader_mock) + return writer_mock, reader_mock + + +# A minimal valid JSON response the server would return. +_VALID_JSON_RESP: dict = { + "status": True, + "is_schema_valid": True, + "is_config_valid": True, + "dataset_valid": True, + "errors": [], + "invalid_datetime_formats": {}, + "dataset_name": "ds", + "dataset_version": 1, + "no_of_columns": 1, + "valid_record_count": 1, + "duplicate_record_count": 0, + "invalid_record_count": 0, +} + + +class TestDryPublishDatasetTransport: + """``ArrowFlightTransport.dry_publish_dataset`` must normalise the Arrow schema, + manage the writer lifecycle correctly, and build ``DryPublishResponse`` from + the server's two-phase metadata response.""" + + # --- schema normalisation --- + + def test_large_string_columns_cast_to_string_before_send(self) -> None: + transport = _make_flight_transport() + _wire_do_put(transport, _VALID_JSON_RESP) + + df = pd.DataFrame({"name": ["Alice", "Bob"], "age": [30, 25]}) + transport.dry_publish_dataset(PublishRequest(input_config="{}", data=df)) + + sent_schema = transport._client.do_put.call_args[0][1] + assert sent_schema.field("name").type == pa.utf8(), "large_string should have been normalised to string" + + def test_integer_and_float_columns_not_affected_by_normalisation(self) -> None: + transport = _make_flight_transport() + _wire_do_put(transport, _VALID_JSON_RESP) + + df = pd.DataFrame({"count": [1, 2], "score": [1.1, 2.2]}) + transport.dry_publish_dataset(PublishRequest(input_config="{}", data=df)) + + sent_schema = transport._client.do_put.call_args[0][1] + assert pa.types.is_integer(sent_schema.field("count").type) + assert pa.types.is_floating(sent_schema.field("score").type) + + # --- writer lifecycle --- + + def test_done_writing_is_called_once_on_success(self) -> None: + transport = _make_flight_transport() + writer_mock, _ = _wire_do_put(transport, _VALID_JSON_RESP) + + transport.dry_publish_dataset(PublishRequest(input_config="{}", data=pd.DataFrame({"x": [1]}))) + writer_mock.done_writing.assert_called_once() + + def test_close_is_called_once_on_success(self) -> None: + transport = _make_flight_transport() + writer_mock, _ = _wire_do_put(transport, _VALID_JSON_RESP) + + transport.dry_publish_dataset(PublishRequest(input_config="{}", data=pd.DataFrame({"x": [1]}))) + writer_mock.close.assert_called_once() + + def test_done_writing_precedes_close(self) -> None: + transport = _make_flight_transport() + writer_mock, _ = _wire_do_put(transport, _VALID_JSON_RESP) + + transport.dry_publish_dataset(PublishRequest(input_config="{}", data=pd.DataFrame({"x": [1]}))) + method_names = [str(c) for c in writer_mock.method_calls] + done_idx = next(i for i, n in enumerate(method_names) if "done_writing" in n) + close_idx = next(i for i, n in enumerate(method_names) if "close" in n) + assert done_idx < close_idx, "done_writing() must be called before close()" + + def test_close_called_even_when_write_batch_raises(self) -> None: + transport = _make_flight_transport() + writer_mock = MagicMock() + writer_mock.write_batch.side_effect = RuntimeError("network error") + reader_mock = MagicMock() + transport._client.do_put.return_value = (writer_mock, reader_mock) + + from dataconnect.transport.errors import TransportError + + with pytest.raises(TransportError): + transport.dry_publish_dataset(PublishRequest(input_config="{}", data=pd.DataFrame({"x": [1]}))) + + writer_mock.close.assert_called_once() + + def test_done_writing_not_called_when_write_batch_raises(self) -> None: + transport = _make_flight_transport() + writer_mock = MagicMock() + writer_mock.write_batch.side_effect = RuntimeError("network error") + transport._client.do_put.return_value = (writer_mock, MagicMock()) + + from dataconnect.transport.errors import TransportError + + with pytest.raises(TransportError): + transport.dry_publish_dataset(PublishRequest(input_config="{}", data=pd.DataFrame({"x": [1]}))) + + writer_mock.done_writing.assert_not_called() + + # --- response parsing --- + + def test_returns_dry_publish_response_instance(self) -> None: + transport = _make_flight_transport() + _wire_do_put(transport, _VALID_JSON_RESP) + + result = transport.dry_publish_dataset(PublishRequest(input_config="{}", data=pd.DataFrame({"x": [1]}))) + assert isinstance(result, DryPublishResponse) + + def test_status_parsed_from_json(self) -> None: + transport = _make_flight_transport() + _wire_do_put(transport, {**_VALID_JSON_RESP, "status": False}) + + result = transport.dry_publish_dataset(PublishRequest(input_config="{}", data=pd.DataFrame({"x": [1]}))) + assert result.status is False + + def test_errors_list_parsed_from_json(self) -> None: + transport = _make_flight_transport() + _wire_do_put(transport, {**_VALID_JSON_RESP, "errors": ["err1", "err2"]}) + + result = transport.dry_publish_dataset(PublishRequest(input_config="{}", data=pd.DataFrame({"x": [1]}))) + assert result.errors == ["err1", "err2"] + + def test_none_ipc_buf_yields_no_invalid_records(self) -> None: + transport = _make_flight_transport() + # Second reader.read() returns None — server sent no invalid-records table + json_buf = _make_json_buf(_VALID_JSON_RESP) + reader_mock = MagicMock() + reader_mock.read.side_effect = [json_buf, None] + writer_mock = MagicMock() + transport._client.do_put.return_value = (writer_mock, reader_mock) + + result = transport.dry_publish_dataset(PublishRequest(input_config="{}", data=pd.DataFrame({"x": [1]}))) + assert result.invalid_records is None + + def test_ipc_buf_is_decoded_to_dataframe(self) -> None: + transport = _make_flight_transport() + invalid_df = pd.DataFrame({"row_id": [10, 20], "reason": ["bad type", "null value"]}) + _wire_do_put(transport, _VALID_JSON_RESP, invalid_records_df=invalid_df) + + result = transport.dry_publish_dataset(PublishRequest(input_config="{}", data=pd.DataFrame({"x": [1]}))) + assert result.invalid_records is not None + assert list(result.invalid_records.columns) == ["row_id", "reason"] + assert len(result.invalid_records) == 2