From de239bc2b8af2b7b2436206fe14e7476a4ec3e06 Mon Sep 17 00:00:00 2001 From: Gaurav Londhe Date: Tue, 29 Sep 2026 22:22:35 +0530 Subject: [PATCH 1/3] feat: Enhance API responses with trace ID --- dataconnect/client.py | 24 ++++--- dataconnect/models.py | 27 +++++++- dataconnect/service/base.py | 7 +- dataconnect/service/default.py | 34 ++++++---- dataconnect/service/mappers.py | 1 + .../transport/arrow_flight/transport.py | 35 +++++++--- dataconnect/transport/base.py | 15 +++- dataconnect/transport/models.py | 21 ++++++ tests/test_dry_publish.py | 30 ++++++-- tests/test_fetch_data.py | 41 ++++++++--- tests/test_get_datasets_paginated.py | 64 +++++++++++++---- tests/test_get_datetime_formats.py | 46 +++++++++---- tests/test_publish.py | 30 ++++++-- tests/test_transport_trace_metadata.py | 68 +++++++++++++++++++ 14 files changed, 357 insertions(+), 86 deletions(-) create mode 100644 tests/test_transport_trace_metadata.py diff --git a/dataconnect/client.py b/dataconnect/client.py index 6acbf22..c4b3c4f 100644 --- a/dataconnect/client.py +++ b/dataconnect/client.py @@ -14,9 +14,10 @@ from dataconnect.models import ( Dataset, - DatasetVersion, + DatasetVersionsResult, DatetimeFormatsResult, DryPublishResult, + FetchDataResult, PaginatedResponse, PublishResult, StudiesResult, @@ -55,19 +56,19 @@ def connect( # Public API def get_studies(self, search_study_name: str | None = None) -> StudiesResult: - """List the studies the client is authorized to access.""" + """List the studies the client is authorized to access, paired with the server's trace id.""" return self._service.get_studies(search_study_name=search_study_name) - def get_dataset_versions(self, dataset_uuid: UUID) -> list[DatasetVersion]: - """List the dataset versions the client is authorized to access.""" + def get_dataset_versions(self, dataset_uuid: UUID) -> DatasetVersionsResult: + """List the dataset versions the client is authorized to access, paired with the server's trace id.""" return self._service.get_dataset_versions(dataset_uuid) def fetch_data( self, dataset_uuid: UUID, first_n_rows: int | None = None, - ) -> pd.DataFrame: - """Fetch data frames for a given dataset UUID.""" + ) -> FetchDataResult: + """Fetch data for a given dataset UUID, paired with the server's trace id.""" return self._service.fetch_data(dataset_uuid, first_n_rows) def get_datasets( @@ -86,7 +87,8 @@ def get_datasets( page_size: Number of results per page. Returns: - A :class:`PaginatedResponse` of :class:`Dataset` items matching the criteria. + A :class:`PaginatedResponse` of :class:`Dataset` items matching the criteria, + with ``trace_id`` set to the server's trace id for this call. """ return self._service.get_datasets( study_environment_uuid=study_environment_uuid, @@ -121,8 +123,8 @@ def dry_publish( Returns: A :class:`DryPublishResult` containing the server's validation - outcome, including per-field validity flags, error messages, and - an optional ``invalid_records`` DataFrame. + outcome, including per-field validity flags, error messages, an + optional ``invalid_records`` DataFrame, and the server's ``trace_id``. """ return self._service.dry_publish( project_token=project_token, @@ -159,7 +161,7 @@ def publish( Returns: A :class:`PublishResult` containing the server's publish outcome, - including dataset UUID, version, and record counts. + including dataset UUID, version, record counts, and the server's ``trace_id``. """ return self._service.publish( project_token=project_token, @@ -189,7 +191,7 @@ def get_datetime_formats( A :class:`DatetimeFormatsResult` exposing the classified list via :meth:`~DatetimeFormatsResult.all` and the type-filtered views via :meth:`~DatetimeFormatsResult.dates` and - :meth:`~DatetimeFormatsResult.datetimes`. + :meth:`~DatetimeFormatsResult.datetimes`, plus the server's ``trace_id``. """ return self._service.get_datetime_formats( project_token=project_token, diff --git a/dataconnect/models.py b/dataconnect/models.py index d96a4e1..0d114eb 100644 --- a/dataconnect/models.py +++ b/dataconnect/models.py @@ -27,6 +27,15 @@ class Study: class StudiesResult: total_records: int studies: list[Study] + trace_id: str | None = None + + +@dataclass(frozen=True) +class FetchDataResult: + """A fetched DataFrame paired with the server's trace id for that call.""" + + data: pd.DataFrame + trace_id: str | None = None @dataclass(frozen=True) @@ -39,20 +48,28 @@ class DatasetVersion: blinding_status: str | None = None +@dataclass(frozen=True) +class DatasetVersionsResult: + """Dataset versions paired with the server's trace id for that call.""" + + items: list[DatasetVersion] + trace_id: str | None = None + + class DatasetFrame: """Lazy dataset reference; fetching requires the originating client to remain open.""" __slots__ = ("_dataset_uuid", "_fetch_data") - def __init__(self, dataset_uuid: str, fetch_data: Callable[[UUID, int | None], pd.DataFrame]) -> None: + def __init__(self, dataset_uuid: str, fetch_data: Callable[[UUID, int | None], FetchDataResult]) -> None: self._dataset_uuid = dataset_uuid self._fetch_data = fetch_data - def head(self, count: int = 6) -> pd.DataFrame: + def head(self, count: int = 6) -> FetchDataResult: """Fetch the first count rows, matching R's default of six rows.""" return self._fetch_data(UUID(self._dataset_uuid), count) - def collect(self) -> pd.DataFrame: + def collect(self) -> FetchDataResult: """Fetch the complete dataset without retaining a previous head limit.""" return self._fetch_data(UUID(self._dataset_uuid), None) @@ -104,6 +121,7 @@ class PaginatedResponse(Generic[T]): # noqa: UP046 total_records: int pagination: Pagination items: list[T] + trace_id: str | None = None @dataclass(frozen=True) @@ -146,6 +164,7 @@ class _PublishEnvelopeResult: checks: ResultChecks = field(default_factory=ResultChecks) errors: list[str] = field(default_factory=list) invalid_records: pd.DataFrame | None = None + trace_id: str | None = None # Flat accessors below are deprecated views onto the envelope, kept so # existing notebooks keep working. Prefer metadata/metrics/checks. @@ -238,6 +257,8 @@ class DatetimeFormatsResult: formats: list[DatetimeFormat] = field(default_factory=list) """The full list of supported formats, in the order returned by the server.""" + trace_id: str | None = None + def all(self) -> list[DatetimeFormat]: """Return every supported format with its type classification.""" return list(self.formats) diff --git a/dataconnect/service/base.py b/dataconnect/service/base.py index 4ac7011..eb06a54 100644 --- a/dataconnect/service/base.py +++ b/dataconnect/service/base.py @@ -9,9 +9,10 @@ from dataconnect.models import ( Dataset, - DatasetVersion, + DatasetVersionsResult, DatetimeFormatsResult, DryPublishResult, + FetchDataResult, PaginatedResponse, PublishResult, StudiesResult, @@ -34,14 +35,14 @@ def get_datasets( ) -> PaginatedResponse[Dataset]: ... @abstractmethod - def get_dataset_versions(self, dataset_uuid: UUID) -> list[DatasetVersion]: ... + def get_dataset_versions(self, dataset_uuid: UUID) -> DatasetVersionsResult: ... @abstractmethod def fetch_data( self, dataset_uuid: UUID, first_n_rows: int | None = None, - ) -> pd.DataFrame: ... + ) -> FetchDataResult: ... @abstractmethod def dry_publish( diff --git a/dataconnect/service/default.py b/dataconnect/service/default.py index 51a2845..e23786d 100644 --- a/dataconnect/service/default.py +++ b/dataconnect/service/default.py @@ -13,10 +13,11 @@ from dataconnect.models import ( Dataset, DatasetFrame, - DatasetVersion, + DatasetVersionsResult, DatetimeFormat, DatetimeFormatsResult, DryPublishResult, + FetchDataResult, PaginatedResponse, Pagination, PublishResult, @@ -92,21 +93,23 @@ def get_studies(self, search_study_name: str | None = None) -> StudiesResult: request = request.append_body({"search_study_name": search_study_name}) try: - resources = self._transport.list_resources(request) + result = self._transport.list_resources(request) + resources = result.resources total_records = resources[0].total_records if resources else 0 studies = [resource_to_study(r) for r in resources] - return StudiesResult(total_records=total_records, studies=studies) + return StudiesResult(total_records=total_records, studies=studies, trace_id=result.trace_id) except Exception as ex: raise translate_error(ex) from ex - def get_dataset_versions(self, dataset_uuid: UUID) -> list[DatasetVersion]: + def get_dataset_versions(self, dataset_uuid: UUID) -> DatasetVersionsResult: """List available versions for a dataset. Args: dataset_uuid: UUID of the dataset whose versions are requested. Returns: - A list of :class:`DatasetVersion` objects for the given dataset. + A :class:`DatasetVersionsResult` containing the dataset versions, + newest first, and the server's trace id for this call. Raises: ValidationError: If *dataset_uuid* is not a valid UUID (upstream). @@ -115,18 +118,19 @@ def get_dataset_versions(self, dataset_uuid: UUID) -> list[DatasetVersion]: request = ResourceQuery(action=_ACTION_LIST_DATASET_VERSIONS).append_body({"dataset_uuid": str(dataset_uuid)}) try: - resources = self._transport.list_resources(request) + result = self._transport.list_resources(request) # Return Sorted dataset versions in descending order (newest first) based on the dataset_version field. - return sorted( - (resource_to_dataset_version(r) for r in resources), + versions = sorted( + (resource_to_dataset_version(r) for r in result.resources), key=lambda dv: dv.dataset_version, reverse=True, ) + return DatasetVersionsResult(items=versions, trace_id=result.trace_id) except Exception as ex: raise translate_error(ex) from ex - def fetch_data(self, dataset_uuid: UUID, first_n_rows: int | None = None) -> pd.DataFrame: + def fetch_data(self, dataset_uuid: UUID, first_n_rows: int | None = None) -> FetchDataResult: """Fetch data for a dataset""" ticket = DatasetTicket( @@ -136,7 +140,7 @@ def fetch_data(self, dataset_uuid: UUID, first_n_rows: int | None = None) -> pd. try: table = self._transport.get_ticket(ticket) - return resource_to_fetched_data(table) + return FetchDataResult(data=resource_to_fetched_data(table), trace_id=table.trace_id) except TransportError as ex: raise translate_error(ex) from ex @@ -169,7 +173,8 @@ def get_datasets( ) try: - resources = self._transport.list_resources(request) + result = self._transport.list_resources(request) + resources = result.resources items = [] for resource in resources: dataset = resource_to_dataset(resource) @@ -183,6 +188,7 @@ def get_datasets( total_records=total_records, pagination=Pagination(page=page, page_size=page_size, total_pages=total_pages), items=items, + trace_id=result.trace_id, ) except TransportError as ex: raise translate_error(ex) from ex @@ -347,7 +353,7 @@ def get_datetime_formats( request = DatetimeFormatsRequest(project_token=project_token, format_type=format_type) try: - raw_formats = self._transport.get_datetime_formats(request) + formats_response = self._transport.get_datetime_formats(request) except TransportError as ex: raise translate_error(ex) from ex @@ -356,10 +362,10 @@ def get_datetime_formats( format=fmt, type="datetime" if "HH:mm" in fmt else "date", ) - for fmt in raw_formats + for fmt in formats_response.formats ] - return DatetimeFormatsResult(formats=formats) + return DatetimeFormatsResult(formats=formats, trace_id=formats_response.trace_id) def close(self) -> None: """Close the underlying transport connection.""" diff --git a/dataconnect/service/mappers.py b/dataconnect/service/mappers.py index 6d56ef9..c5619bb 100644 --- a/dataconnect/service/mappers.py +++ b/dataconnect/service/mappers.py @@ -147,6 +147,7 @@ def _envelope_to_domain(envelope: PublishEnvelope, result_cls: type[_ResultT]) - ), errors=envelope.errors, invalid_records=envelope.invalid_records, + trace_id=envelope.trace_id, ) diff --git a/dataconnect/transport/arrow_flight/transport.py b/dataconnect/transport/arrow_flight/transport.py index 7dfe1ff..39e495e 100644 --- a/dataconnect/transport/arrow_flight/transport.py +++ b/dataconnect/transport/arrow_flight/transport.py @@ -26,10 +26,12 @@ DatasetTicket, DataTable, DatetimeFormatsRequest, + DatetimeFormatsResponse, DryPublishResponse, PublishRequest, PublishResponse, ResourceInfo, + ResourceListResult, ResourceQuery, ) @@ -49,7 +51,7 @@ def _to_resource_info(info: flight.FlightInfo) -> ResourceInfo: ) -def _to_bytes(table: pa.Table) -> DataTable: +def _to_bytes(table: pa.Table, trace_id: str | None = None) -> DataTable: """Serialize a ``pa.Table`` to a technology-agnostic ``DataTable``. Each record batch is serialized individually as Arrow IPC bytes. @@ -65,7 +67,7 @@ def _to_bytes(table: pa.Table) -> DataTable: writer.close() ipc_bytes = sink.getvalue().to_pybytes() - return DataTable(schema_bytes=schema_bytes, ipc_bytes=ipc_bytes) + return DataTable(schema_bytes=schema_bytes, ipc_bytes=ipc_bytes, trace_id=trace_id) def _normalize_arrow_type(dtype: pa.DataType) -> pa.DataType: @@ -184,7 +186,7 @@ def _options(self) -> flight.FlightCallOptions: # Transport - def list_resources(self, request: ResourceQuery) -> list[ResourceInfo]: + def list_resources(self, request: ResourceQuery) -> ResourceListResult: """Translate the action name to Arrow Flight criteria and return resource records.""" flight_type = _ACTION_FLIGHT_TYPE.get(request.action) @@ -200,8 +202,16 @@ def list_resources(self, request: ResourceQuery) -> list[ResourceInfo]: criteria = json.dumps({**body, "flight_type": flight_type}, separators=(",", ":")).encode("utf-8") try: - raw_flights = self._client.list_flights(criteria, self._options()) - return [_to_resource_info(f) for f in raw_flights] + raw_flights = list(self._client.list_flights(criteria, self._options())) + + trace_id = None + if raw_flights and raw_flights[0].app_metadata: + try: + trace_id = json.loads(raw_flights[0].app_metadata.decode("utf-8")).get("trace_id") + except (json.JSONDecodeError, UnicodeDecodeError): + trace_id = None + + return ResourceListResult(resources=[_to_resource_info(f) for f in raw_flights], trace_id=trace_id) except Exception as ex: raise parse_dataconnect_error(ex) from ex @@ -224,7 +234,13 @@ def get_ticket(self, ticket: DatasetTicket) -> DataTable: except flight.FlightError as ex: raise parse_dataconnect_error(ex) from ex - return _to_bytes(pa.Table.from_batches(batches, schema=table.schema)) + trace_id = None + if table.schema.metadata: + raw_trace_id = table.schema.metadata.get(b"trace_id") + if raw_trace_id is not None: + trace_id = raw_trace_id.decode("utf-8") + + return _to_bytes(pa.Table.from_batches(batches, schema=table.schema), trace_id=trace_id) except Exception as ex: raise parse_dataconnect_error(ex) from ex @@ -355,7 +371,7 @@ def publish_dataset(self, publish_request: PublishRequest) -> PublishResponse: except Exception as ex: raise parse_dataconnect_error(ex) from ex - def get_datetime_formats(self, request: DatetimeFormatsRequest) -> list[str]: + def get_datetime_formats(self, request: DatetimeFormatsRequest) -> DatetimeFormatsResponse: """Invoke the Arrow Flight ``get_datetime_formats`` action and return the format list. The server expects a JSON body of the form @@ -382,9 +398,10 @@ def get_datetime_formats(self, request: DatetimeFormatsRequest) -> list[str]: results = list(self._client.do_action(action, self._options())) if not results: - return [] + return DatetimeFormatsResponse(formats=[], trace_id=None) - return json.loads(results[0].body.to_pybytes().decode("utf-8")) + response = json.loads(results[0].body.to_pybytes().decode("utf-8")) + return DatetimeFormatsResponse(formats=response.get("formats", []), trace_id=response.get("trace_id")) except Exception as ex: raise parse_dataconnect_error(ex) from ex diff --git a/dataconnect/transport/base.py b/dataconnect/transport/base.py index 91ce61f..956c828 100644 --- a/dataconnect/transport/base.py +++ b/dataconnect/transport/base.py @@ -13,10 +13,11 @@ DatasetTicket, DataTable, DatetimeFormatsRequest, + DatetimeFormatsResponse, DryPublishResponse, PublishRequest, PublishResponse, - ResourceInfo, + ResourceListResult, ResourceQuery, ) @@ -25,11 +26,15 @@ class Transport(ABC): """Minimal abstract transport for DataConnect operations.""" @abstractmethod - def list_resources(self, request: ResourceQuery) -> list[ResourceInfo]: + def list_resources(self, request: ResourceQuery) -> ResourceListResult: """List available data resources matching the given query. The transport does not interpret the action name or body — that is the service layer's responsibility. + + Returns: + A :class:`ResourceListResult` with the matched resources and the + server's trace id for this call, or ``None`` if not provided. """ @abstractmethod @@ -57,12 +62,16 @@ def publish_dataset(self, publish_request: PublishRequest) -> PublishResponse: """ @abstractmethod - def get_datetime_formats(self, request: DatetimeFormatsRequest) -> list[str]: + def get_datetime_formats(self, request: DatetimeFormatsRequest) -> DatetimeFormatsResponse: """Return the supported datetime format strings for the project. The transport does not interpret ``format_type`` — that is the service layer's responsibility. The server applies the filter and returns the already-filtered list of format strings. + + Returns: + A :class:`DatetimeFormatsResponse` with the filtered formats and + the server's trace id for this call, or ``None`` if not provided. """ @abstractmethod diff --git a/dataconnect/transport/models.py b/dataconnect/transport/models.py index 5befc2d..a389a83 100644 --- a/dataconnect/transport/models.py +++ b/dataconnect/transport/models.py @@ -53,16 +53,26 @@ class ResourceInfo: schema_bytes: bytes +@dataclass(frozen=True) +class ResourceListResult: + """Result of a resource-listing call (studies/datasets/dataset_versions).""" + + resources: list[ResourceInfo] + trace_id: str | None = None + + @dataclass(frozen=True) class DataTable: """Technology-agnostic representation of a fetched data result. ``schema_bytes`` holds the Arrow IPC-serialized schema. ``ipc_bytes`` holds the full Arrow IPC stream (schema + all batches). + ``trace_id`` is the server's trace id for the call that produced this result, if any. """ schema_bytes: bytes ipc_bytes: bytes + trace_id: str | None = None @dataclass(frozen=True) @@ -80,6 +90,14 @@ class DatetimeFormatsRequest: format_type: str = "all" +@dataclass(frozen=True) +class DatetimeFormatsResponse: + """Result of a get_datetime_formats call.""" + + formats: list[str] + trace_id: str | None = None + + @dataclass(frozen=True) class PublishRequest: """A publish request containing the input configuration and the dataset to be published. @@ -143,6 +161,8 @@ class PublishEnvelope: errors: list[str] = field(default_factory=list) invalid_records: pd.DataFrame | None = None """Populated from the Arrow IPC channel, not from the JSON payload.""" + trace_id: str | None = None + """Populated from the JSON payload.""" @classmethod def from_json(cls, payload: dict, invalid_records: pd.DataFrame | None = None) -> PublishEnvelope: @@ -178,6 +198,7 @@ def from_json(cls, payload: dict, invalid_records: pd.DataFrame | None = None) - ), errors=payload.get("errors") or [], invalid_records=invalid_records, + trace_id=payload.get("trace_id"), ) diff --git a/tests/test_dry_publish.py b/tests/test_dry_publish.py index 45d4dd6..f8210c3 100644 --- a/tests/test_dry_publish.py +++ b/tests/test_dry_publish.py @@ -30,10 +30,11 @@ from dataconnect.transport.models import ( DatasetTicket, DataTable, + DatetimeFormatsResponse, DryPublishResponse, PublishRequest, PublishResponse, - ResourceInfo, + ResourceListResult, ResourceQuery, ResponseChecks, ResponseMetadata, @@ -84,6 +85,7 @@ def _make_dry_publish_response(**overrides: object) -> DryPublishResponse: ), errors=flat["errors"], invalid_records=flat["invalid_records"], + trace_id=overrides.get("trace_id"), # type: ignore[arg-type] ) @@ -225,6 +227,14 @@ def test_returns_dry_publish_result_instance(self) -> None: result = dry_publish_response_to_domain(_make_dry_publish_response()) assert isinstance(result, DryPublishResult) + def test_trace_id_mapped(self) -> None: + result = dry_publish_response_to_domain(_make_dry_publish_response(trace_id="trace-dry-1")) + assert result.trace_id == "trace-dry-1" + + def test_trace_id_defaults_to_none(self) -> None: + result = dry_publish_response_to_domain(_make_dry_publish_response()) + assert result.trace_id is None + # --------------------------------------------------------------------------- # DefaultDataConnectService.dry_publish @@ -243,8 +253,8 @@ def __init__( self._raise = raise_error self.last_request: PublishRequest | None = None - def list_resources(self, request: ResourceQuery) -> list[ResourceInfo]: - return [] + def list_resources(self, request: ResourceQuery) -> ResourceListResult: + return ResourceListResult(resources=[], trace_id=None) def get_ticket(self, ticket: DatasetTicket) -> DataTable: raise NotImplementedError @@ -258,7 +268,7 @@ def dry_publish_dataset(self, publish_request: PublishRequest) -> DryPublishResp def publish_dataset(self, publish_request: PublishRequest) -> PublishResponse: raise NotImplementedError - def get_datetime_formats(self, request: DatetimeFormatsResult) -> list[str]: # type: ignore[override] + def get_datetime_formats(self, request: DatetimeFormatsResult) -> DatetimeFormatsResponse: # type: ignore[override] raise NotImplementedError def close(self) -> None: @@ -292,6 +302,11 @@ def test_returns_dry_publish_result_instance(self) -> None: result = service.dry_publish(**_default_dry_publish_args()) assert isinstance(result, DryPublishResult) + def test_trace_id_is_propagated_from_transport(self) -> None: + service, _ = _make_service(dry_publish_return=_make_dry_publish_response(trace_id="trace-dry-svc-1")) + result = service.dry_publish(**_default_dry_publish_args()) + assert result.trace_id == "trace-dry-svc-1" + 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 @@ -514,6 +529,13 @@ def test_returns_dry_publish_response_instance(self) -> None: result = transport.dry_publish_dataset(PublishRequest(input_config="{}", data=pd.DataFrame({"x": [1]}))) assert isinstance(result, DryPublishResponse) + def test_trace_id_is_read_from_json_payload(self) -> None: + transport = _make_flight_transport() + _wire_do_put(transport, {**_VALID_JSON_RESP, "trace_id": "trace-put-1"}) + + result = transport.dry_publish_dataset(PublishRequest(input_config="{}", data=pd.DataFrame({"x": [1]}))) + assert result.trace_id == "trace-put-1" + def test_status_parsed_from_json(self) -> None: transport = _make_flight_transport() _wire_do_put(transport, {**_VALID_JSON_RESP, "success": False}) diff --git a/tests/test_fetch_data.py b/tests/test_fetch_data.py index a431386..f1d5a1b 100644 --- a/tests/test_fetch_data.py +++ b/tests/test_fetch_data.py @@ -22,7 +22,7 @@ TransportNotFoundError, TransportServerError, ) -from dataconnect.transport.models import DatasetTicket, DataTable, ResourceInfo, ResourceQuery +from dataconnect.transport.models import DatasetTicket, DataTable, ResourceListResult, ResourceQuery # --------------------------------------------------------------------------- # Fake transport @@ -41,8 +41,8 @@ def __init__( self._get_ticket_error = get_ticket_error self.last_ticket: DatasetTicket | None = None - def list_resources(self, request: ResourceQuery) -> list[ResourceInfo]: - return [] + def list_resources(self, request: ResourceQuery) -> ResourceListResult: + return ResourceListResult(resources=[], trace_id=None) def get_ticket(self, ticket: DatasetTicket) -> DataTable: self.last_ticket = ticket @@ -55,7 +55,7 @@ def close(self) -> None: return None -def _make_ipc_table(data: dict) -> DataTable: +def _make_ipc_table(data: dict, trace_id: str | None = None) -> DataTable: arrow_table = pa.table(data) sink = pa.BufferOutputStream() writer = pa.ipc.new_stream(sink, arrow_table.schema) @@ -64,6 +64,7 @@ def _make_ipc_table(data: dict) -> DataTable: return DataTable( schema_bytes=arrow_table.schema.serialize().to_pybytes(), ipc_bytes=sink.getvalue().to_pybytes(), + trace_id=trace_id, ) @@ -80,9 +81,9 @@ def test_fetch_data_returns_dataframe_with_correct_values() -> None: 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"] + assert isinstance(result.data, pd.DataFrame) + assert result.data["subject_id"].tolist() == source["subject_id"] + assert result.data["age"].tolist() == source["age"] def test_fetch_data_builds_correct_ticket() -> None: @@ -97,6 +98,26 @@ def test_fetch_data_builds_correct_ticket() -> None: assert transport.last_ticket.limit == 10 +def test_fetch_data_returns_trace_id_from_transport() -> None: + dataset_uuid = UUID("073410b6-79be-3e7d-ae37-92f6e054013e") + transport = _FakeTransport(data_table=_make_ipc_table({"x": [1]}, trace_id="abc123")) + service = DefaultDataConnectService(transport) + + result = service.fetch_data(dataset_uuid) + + assert result.trace_id == "abc123" + + +def test_fetch_data_trace_id_defaults_to_none() -> None: + dataset_uuid = UUID("073410b6-79be-3e7d-ae37-92f6e054013e") + transport = _FakeTransport(data_table=_make_ipc_table({"x": [1]})) + service = DefaultDataConnectService(transport) + + result = service.fetch_data(dataset_uuid) + + assert result.trace_id is None + + 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]})) @@ -147,9 +168,9 @@ def test_fetch_data_returns_empty_dataframe_for_empty_table() -> None: result = service.fetch_data(dataset_uuid) - assert isinstance(result, pd.DataFrame) - assert len(result) == 0 - assert "col" in result.columns + assert isinstance(result.data, pd.DataFrame) + assert len(result.data) == 0 + assert "col" in result.data.columns def test_fetch_data_translates_authentication_error() -> None: diff --git a/tests/test_get_datasets_paginated.py b/tests/test_get_datasets_paginated.py index d3cae99..bad16a8 100644 --- a/tests/test_get_datasets_paginated.py +++ b/tests/test_get_datasets_paginated.py @@ -16,7 +16,14 @@ from dataconnect.models import Dataset, PaginatedResponse, Pagination from dataconnect.service.default import DefaultDataConnectService from dataconnect.transport.errors import TransportError -from dataconnect.transport.models import DataRef, DatasetTicket, DataTable, ResourceInfo, ResourceQuery +from dataconnect.transport.models import ( + DataRef, + DatasetTicket, + DataTable, + ResourceInfo, + ResourceListResult, + ResourceQuery, +) class _FakeTransport: @@ -24,16 +31,18 @@ def __init__( self, resources: list[ResourceInfo] | None = None, error: Exception | None = None, + trace_id: str | None = None, ) -> None: self._resources = resources or [] self._error = error + self._trace_id = trace_id self.last_request: ResourceQuery | None = None - def list_resources(self, request: ResourceQuery) -> list[ResourceInfo]: + def list_resources(self, request: ResourceQuery) -> ResourceListResult: self.last_request = request if self._error is not None: raise self._error - return self._resources + return ResourceListResult(resources=self._resources, trace_id=self._trace_id) def do_get(self, request: ResourceQuery) -> None: pass @@ -156,6 +165,22 @@ def test_translates_transport_errors(self) -> None: with pytest.raises(Exception, match="cannot connect"): service.get_datasets(study_environment_uuid=_STUDY_ENV_UUID) + def test_returns_trace_id_from_transport(self) -> None: + transport = _FakeTransport(resources=[], trace_id="trace-datasets-1") + service = DefaultDataConnectService(transport) + + result = service.get_datasets(study_environment_uuid=_STUDY_ENV_UUID) + + assert result.trace_id == "trace-datasets-1" + + def test_trace_id_defaults_to_none(self) -> None: + transport = _FakeTransport(resources=[]) + service = DefaultDataConnectService(transport) + + result = service.get_datasets(study_environment_uuid=_STUDY_ENV_UUID) + + assert result.trace_id is None + def test_multiple_items_returned(self) -> None: resources = [ _dataset_resource( @@ -227,11 +252,12 @@ def __init__(self, pages: dict[int, list[dict[str, object]]]) -> None: self.lock = Lock() self.fetch_error: TransportError | None = None - def list_resources(self, request: ResourceQuery) -> list[ResourceInfo]: + def list_resources(self, request: ResourceQuery) -> ResourceListResult: self.requests.append(request) body = json.loads(request.body) total_records = sum(len(payloads) for payloads in self.pages.values()) - return [_dataset_resource(payload, total_records) for payload in self.pages[body["page"]]] + resources = [_dataset_resource(payload, total_records) for payload in self.pages[body["page"]]] + return ResourceListResult(resources=resources, trace_id=None) def get_ticket(self, ticket: DatasetTicket) -> DataTable: if self.closed: @@ -331,13 +357,13 @@ def test_frame_fetches_only_on_demand_and_binds_each_dataset() -> None: assert second.frame is not None assert first.frame is not second.frame - preview = first.frame.head(3) + preview = first.frame.head(3).data assert isinstance(preview, pd.DataFrame) assert preview["value"].tolist() == [0, 1, 2] assert preview["dataset_uuid"].tolist() == [_DATASET_UUID] * 3 - assert len(first.frame.head()) == 6 - assert len(first.frame.collect()) == 8 - assert second.frame.head(1)["dataset_uuid"].tolist() == [_OTHER_DATASET_UUID] + assert len(first.frame.head().data) == 6 + assert len(first.frame.collect().data) == 8 + assert second.frame.head(1).data["dataset_uuid"].tolist() == [_OTHER_DATASET_UUID] assert transport.tickets == [ DatasetTicket(dataset_uuid=_DATASET_UUID, limit=3), DatasetTicket(dataset_uuid=_DATASET_UUID, limit=6), @@ -361,7 +387,7 @@ def test_frame_copy_and_inspection_do_not_copy_or_fetch_connection() -> None: for name, value in _METADATA.items(): assert encoded[name] == value assert transport.tickets == [] - assert len(encoded["frame"].head(1)) == 1 + assert len(encoded["frame"].head(1).data) == 1 def test_dataset_constructor_and_equality_remain_independent_of_frame() -> None: @@ -468,12 +494,24 @@ def test_dataset_versions_response_is_unchanged_before_and_after_listing() -> No ] listed = client.get_dataset_versions(UUID(_DATASET_UUID)) - assert [asdict(item) for item in listed] == expected - assert listed[0].blinding_status is None + assert [asdict(item) for item in listed.items] == expected + assert listed.items[0].blinding_status is None transport._resources = [_dataset_resource({**_IDENTIFIERS, **_METADATA})] assert client.get_datasets(_STUDY_ENV_UUID).items[0].frame is not None transport._resources = versions - assert [asdict(item) for item in client.get_dataset_versions(UUID(_DATASET_UUID))] == expected + assert [asdict(item) for item in client.get_dataset_versions(UUID(_DATASET_UUID)).items] == expected assert transport.last_request is not None assert transport.last_request.action == "dataset_versions.list" assert json.loads(transport.last_request.body) == {"dataset_uuid": _DATASET_UUID} + + +def test_get_dataset_versions_returns_trace_id_from_transport() -> None: + transport = _FakeTransport( + resources=[_dataset_resource({**_IDENTIFIERS, "dataset_version": "1"})], + trace_id="trace-versions-1", + ) + service = DefaultDataConnectService(transport) + + result = service.get_dataset_versions(UUID(_DATASET_UUID)) + + assert result.trace_id == "trace-versions-1" diff --git a/tests/test_get_datetime_formats.py b/tests/test_get_datetime_formats.py index 5a754c2..74df91a 100644 --- a/tests/test_get_datetime_formats.py +++ b/tests/test_get_datetime_formats.py @@ -25,10 +25,11 @@ DatasetTicket, DataTable, DatetimeFormatsRequest, + DatetimeFormatsResponse, DryPublishResponse, PublishRequest, PublishResponse, - ResourceInfo, + ResourceListResult, ResourceQuery, ) @@ -52,13 +53,15 @@ def __init__( self, formats: list[str] | None = None, raise_error: Exception | None = None, + trace_id: str | None = None, ) -> None: self._formats = formats if formats is not None else list(_SAMPLE_FORMATS) self._raise = raise_error + self._trace_id = trace_id self.last_request: DatetimeFormatsRequest | None = None - def list_resources(self, request: ResourceQuery) -> list[ResourceInfo]: - return [] + def list_resources(self, request: ResourceQuery) -> ResourceListResult: + return ResourceListResult(resources=[], trace_id=None) def get_ticket(self, ticket: DatasetTicket) -> DataTable: raise NotImplementedError @@ -69,11 +72,11 @@ def dry_publish_dataset(self, publish_request: PublishRequest) -> DryPublishResp def publish_dataset(self, publish_request: PublishRequest) -> PublishResponse: raise NotImplementedError - def get_datetime_formats(self, request: DatetimeFormatsRequest) -> list[str]: + def get_datetime_formats(self, request: DatetimeFormatsRequest) -> DatetimeFormatsResponse: self.last_request = request if self._raise is not None: raise self._raise - return self._formats + return DatetimeFormatsResponse(formats=self._formats, trace_id=self._trace_id) def close(self) -> None: pass @@ -82,8 +85,9 @@ def close(self) -> None: def _make_service( formats: list[str] | None = None, raise_error: Exception | None = None, + trace_id: str | None = None, ) -> tuple[DefaultDataConnectService, _StubTransport]: - transport = _StubTransport(formats=formats, raise_error=raise_error) + transport = _StubTransport(formats=formats, raise_error=raise_error, trace_id=trace_id) return DefaultDataConnectService(transport), transport @@ -249,6 +253,16 @@ def test_empty_server_response_yields_empty_result(self) -> None: assert result.dates() == [] assert result.datetimes() == [] + def test_trace_id_is_propagated_from_transport(self) -> None: + service, _ = _make_service(trace_id="trace-fmt-1") + result = service.get_datetime_formats(project_token="tok") + assert result.trace_id == "trace-fmt-1" + + def test_trace_id_defaults_to_none(self) -> None: + service, _ = _make_service() + result = service.get_datetime_formats(project_token="tok") + assert result.trace_id is None + # --- error translation --- def test_transport_validation_error_is_translated_to_service_error(self) -> None: @@ -274,9 +288,9 @@ def _make_flight_transport() -> ArrowFlightTransport: return ArrowFlightTransport(host="localhost", port=5005, use_tls=False) -def _make_action_result(payload: list[str]) -> MagicMock: +def _make_action_result(formats: list[str], trace_id: str | None = None) -> MagicMock: body = MagicMock() - body.to_pybytes.return_value = json.dumps(payload).encode("utf-8") + body.to_pybytes.return_value = json.dumps({"formats": formats, "trace_id": trace_id}).encode("utf-8") result = MagicMock() result.body = body return result @@ -310,17 +324,25 @@ def test_returns_decoded_format_list(self) -> None: transport = _make_flight_transport() transport._client.do_action.return_value = iter([_make_action_result(_SAMPLE_FORMATS)]) - result = transport.get_datetime_formats(DatetimeFormatsRequest(project_token="tok", format_type="all")) + response = transport.get_datetime_formats(DatetimeFormatsRequest(project_token="tok", format_type="all")) - assert result == _SAMPLE_FORMATS + assert response.formats == _SAMPLE_FORMATS def test_empty_result_iterator_returns_empty_list(self) -> None: transport = _make_flight_transport() transport._client.do_action.return_value = iter([]) - result = transport.get_datetime_formats(DatetimeFormatsRequest(project_token="tok", format_type="all")) + response = transport.get_datetime_formats(DatetimeFormatsRequest(project_token="tok", format_type="all")) + + assert response.formats == [] + + def test_trace_id_is_read_from_json_payload(self) -> None: + transport = _make_flight_transport() + transport._client.do_action.return_value = iter([_make_action_result(_SAMPLE_FORMATS, trace_id="trace-fmt-2")]) + + response = transport.get_datetime_formats(DatetimeFormatsRequest(project_token="tok", format_type="all")) - assert result == [] + assert response.trace_id == "trace-fmt-2" def test_underlying_exception_is_translated_to_transport_error(self) -> None: from dataconnect.transport.errors import TransportError diff --git a/tests/test_publish.py b/tests/test_publish.py index 35e531d..2f8ea10 100644 --- a/tests/test_publish.py +++ b/tests/test_publish.py @@ -30,10 +30,11 @@ DatasetTicket, DataTable, DatetimeFormatsRequest, + DatetimeFormatsResponse, DryPublishResponse, PublishRequest, PublishResponse, - ResourceInfo, + ResourceListResult, ResourceQuery, ResponseMetadata, ResponseMetrics, @@ -72,6 +73,7 @@ def _make_publish_response(**overrides: object) -> PublishResponse: total_duplicate_rows=flat["duplicate_record_count"], ), invalid_records=flat["invalid_records"], + trace_id=overrides.get("trace_id"), # type: ignore[arg-type] ) @@ -146,6 +148,14 @@ def test_returns_publish_result_instance(self) -> None: result = publish_response_to_domain(_make_publish_response()) assert isinstance(result, PublishResult) + def test_trace_id_mapped(self) -> None: + result = publish_response_to_domain(_make_publish_response(trace_id="trace-pub-1")) + assert result.trace_id == "trace-pub-1" + + def test_trace_id_defaults_to_none(self) -> None: + result = publish_response_to_domain(_make_publish_response()) + assert result.trace_id is None + # --------------------------------------------------------------------------- # DefaultDataConnectService.publish @@ -164,8 +174,8 @@ def __init__( self._raise = raise_error self.last_request: PublishRequest | None = None - def list_resources(self, request: ResourceQuery) -> list[ResourceInfo]: - return [] + def list_resources(self, request: ResourceQuery) -> ResourceListResult: + return ResourceListResult(resources=[], trace_id=None) def get_ticket(self, ticket: DatasetTicket) -> DataTable: raise NotImplementedError @@ -179,7 +189,7 @@ def publish_dataset(self, publish_request: PublishRequest) -> PublishResponse: raise self._raise return self._return # type: ignore[return-value] - def get_datetime_formats(self, request: DatetimeFormatsRequest) -> list[str]: # type: ignore[override] + def get_datetime_formats(self, request: DatetimeFormatsRequest) -> DatetimeFormatsResponse: # type: ignore[override] raise NotImplementedError def close(self) -> None: @@ -213,6 +223,11 @@ def test_returns_publish_result_instance(self) -> None: result = service.publish(**_default_publish_args()) assert isinstance(result, PublishResult) + def test_trace_id_is_propagated_from_transport(self) -> None: + service, _ = _make_service(publish_return=_make_publish_response(trace_id="trace-pub-svc-1")) + result = service.publish(**_default_publish_args()) + assert result.trace_id == "trace-pub-svc-1" + def test_successful_status_is_propagated(self) -> None: service, _ = _make_service(publish_return=_make_publish_response(status=True)) assert service.publish(**_default_publish_args()).status is True @@ -435,6 +450,13 @@ def test_returns_publish_response_instance(self) -> None: result = transport.publish_dataset(PublishRequest(input_config="{}", data=pd.DataFrame({"x": [1]}))) assert isinstance(result, PublishResponse) + def test_trace_id_is_read_from_json_payload(self) -> None: + transport = _make_flight_transport() + _wire_do_put(transport, {**_VALID_JSON_RESP, "trace_id": "trace-put-2"}) + + result = transport.publish_dataset(PublishRequest(input_config="{}", data=pd.DataFrame({"x": [1]}))) + assert result.trace_id == "trace-put-2" + def test_status_parsed_from_json(self) -> None: transport = _make_flight_transport() _wire_do_put(transport, {**_VALID_JSON_RESP, "success": False}) diff --git a/tests/test_transport_trace_metadata.py b/tests/test_transport_trace_metadata.py new file mode 100644 index 0000000..b7d9e02 --- /dev/null +++ b/tests/test_transport_trace_metadata.py @@ -0,0 +1,68 @@ +"""Tests for trace metadata decoding in the concrete Arrow Flight transport.""" + +from unittest.mock import MagicMock, patch + +import pyarrow as pa +import pytest + +from dataconnect.transport.arrow_flight.transport import ArrowFlightTransport +from dataconnect.transport.models import DatasetTicket, ResourceQuery + + +def _make_transport() -> ArrowFlightTransport: + with patch.object(ArrowFlightTransport, "_get_client", return_value=MagicMock()): + return ArrowFlightTransport(host="localhost", port=5005, use_tls=False) + + +def _make_flight_info(app_metadata: bytes | None) -> MagicMock: + info = MagicMock() + info.app_metadata = app_metadata + info.descriptor = None + info.endpoints = [] + info.schema = pa.schema([]) + info.total_records = 0 + return info + + +def test_list_resources_reads_trace_id_from_app_metadata() -> None: + transport = _make_transport() + transport._client.list_flights.return_value = [_make_flight_info(b'{"trace_id":"trace-list-1"}')] + + result = transport.list_resources(ResourceQuery(action="studies.list")) + + assert result.trace_id == "trace-list-1" + assert len(result.resources) == 1 + + +@pytest.mark.parametrize("app_metadata", [None, b"", b"{}", b"not-json", b"\xff"]) +def test_list_resources_returns_no_trace_id_for_missing_or_malformed_metadata(app_metadata: bytes | None) -> None: + transport = _make_transport() + transport._client.list_flights.return_value = [_make_flight_info(app_metadata)] + + result = transport.list_resources(ResourceQuery(action="studies.list")) + + assert result.trace_id is None + + +def test_get_ticket_reads_trace_id_from_schema_metadata() -> None: + transport = _make_transport() + reader = MagicMock() + reader.schema = pa.schema([("id", pa.int64())], metadata={b"trace_id": b"trace-get-1"}) + reader.read_chunk.side_effect = StopIteration + transport._client.do_get.return_value = reader + + result = transport.get_ticket(DatasetTicket(dataset_uuid="dataset-1")) + + assert result.trace_id == "trace-get-1" + + +def test_get_ticket_returns_no_trace_id_when_schema_metadata_is_absent() -> None: + transport = _make_transport() + reader = MagicMock() + reader.schema = pa.schema([("id", pa.int64())]) + reader.read_chunk.side_effect = StopIteration + transport._client.do_get.return_value = reader + + result = transport.get_ticket(DatasetTicket(dataset_uuid="dataset-1")) + + assert result.trace_id is None From 5027490df3c726c874f0d646ba54ae88e2b5ef5b Mon Sep 17 00:00:00 2001 From: Gaurav Londhe Date: Wed, 30 Sep 2026 14:08:23 +0530 Subject: [PATCH 2/3] Refactor DataConnect service and transport layers to enhance traceability and simplify response structures --- dataconnect/client.py | 27 ++-- dataconnect/exceptions.py | 5 + dataconnect/models.py | 27 +--- dataconnect/service/base.py | 7 +- dataconnect/service/default.py | 31 ++--- dataconnect/service/error_handler.py | 40 +++++- dataconnect/service/mappers.py | 1 - .../transport/arrow_flight/error_handler.py | 7 ++ .../transport/arrow_flight/transport.py | 119 +++++++++++++++--- dataconnect/transport/base.py | 13 +- dataconnect/transport/errors.py | 1 + dataconnect/transport/models.py | 21 ---- tests/test_dry_publish.py | 33 ++--- tests/test_fetch_data.py | 41 ++---- tests/test_get_datasets_paginated.py | 39 +++--- tests/test_get_datetime_formats.py | 36 ++---- tests/test_publish.py | 30 +---- tests/test_transport_trace_metadata.py | 58 ++++++++- 18 files changed, 300 insertions(+), 236 deletions(-) diff --git a/dataconnect/client.py b/dataconnect/client.py index c4b3c4f..a26c151 100644 --- a/dataconnect/client.py +++ b/dataconnect/client.py @@ -14,10 +14,9 @@ from dataconnect.models import ( Dataset, - DatasetVersionsResult, + DatasetVersion, DatetimeFormatsResult, DryPublishResult, - FetchDataResult, PaginatedResponse, PublishResult, StudiesResult, @@ -35,6 +34,11 @@ def __init__(self, service: DataConnectService) -> None: """Initialize the client with an injected service implementation.""" self._service = service + @property + def trace_id(self) -> str | None: + """Trace ID from the most recent sequential request, when available.""" + return getattr(self._service, "trace_id", None) + @classmethod def connect( cls, @@ -56,19 +60,19 @@ def connect( # Public API def get_studies(self, search_study_name: str | None = None) -> StudiesResult: - """List the studies the client is authorized to access, paired with the server's trace id.""" + """List the studies the client is authorized to access.""" return self._service.get_studies(search_study_name=search_study_name) - def get_dataset_versions(self, dataset_uuid: UUID) -> DatasetVersionsResult: - """List the dataset versions the client is authorized to access, paired with the server's trace id.""" + 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 fetch_data( self, dataset_uuid: UUID, first_n_rows: int | None = None, - ) -> FetchDataResult: - """Fetch data for a given dataset UUID, paired with the server's trace id.""" + ) -> pd.DataFrame: + """Fetch data for a given dataset UUID.""" return self._service.fetch_data(dataset_uuid, first_n_rows) def get_datasets( @@ -87,8 +91,7 @@ def get_datasets( page_size: Number of results per page. Returns: - A :class:`PaginatedResponse` of :class:`Dataset` items matching the criteria, - with ``trace_id`` set to the server's trace id for this call. + A :class:`PaginatedResponse` of :class:`Dataset` items matching the criteria. """ return self._service.get_datasets( study_environment_uuid=study_environment_uuid, @@ -124,7 +127,7 @@ def dry_publish( Returns: A :class:`DryPublishResult` containing the server's validation outcome, including per-field validity flags, error messages, an - optional ``invalid_records`` DataFrame, and the server's ``trace_id``. + optional ``invalid_records`` DataFrame. """ return self._service.dry_publish( project_token=project_token, @@ -161,7 +164,7 @@ def publish( Returns: A :class:`PublishResult` containing the server's publish outcome, - including dataset UUID, version, record counts, and the server's ``trace_id``. + including dataset UUID, version, and record counts. """ return self._service.publish( project_token=project_token, @@ -191,7 +194,7 @@ def get_datetime_formats( A :class:`DatetimeFormatsResult` exposing the classified list via :meth:`~DatetimeFormatsResult.all` and the type-filtered views via :meth:`~DatetimeFormatsResult.dates` and - :meth:`~DatetimeFormatsResult.datetimes`, plus the server's ``trace_id``. + :meth:`~DatetimeFormatsResult.datetimes`. """ return self._service.get_datetime_formats( project_token=project_token, diff --git a/dataconnect/exceptions.py b/dataconnect/exceptions.py index 4329407..9bede5f 100644 --- a/dataconnect/exceptions.py +++ b/dataconnect/exceptions.py @@ -48,6 +48,7 @@ class DataConnectError(Exception): message: str timestamp: str | None = None details: list[ErrorDetail] | None = None + trace_id: str | None = None def __str__(self) -> str: lines = [ @@ -63,6 +64,10 @@ def __str__(self) -> str: for detail in self.details: lines.append(str(detail)) + if self.trace_id is not None: + indent = " " if self.details else "" + lines.append(f"{indent}Trace ID: {self.trace_id}") + return "\n".join(lines) diff --git a/dataconnect/models.py b/dataconnect/models.py index 0d114eb..d96a4e1 100644 --- a/dataconnect/models.py +++ b/dataconnect/models.py @@ -27,15 +27,6 @@ class Study: class StudiesResult: total_records: int studies: list[Study] - trace_id: str | None = None - - -@dataclass(frozen=True) -class FetchDataResult: - """A fetched DataFrame paired with the server's trace id for that call.""" - - data: pd.DataFrame - trace_id: str | None = None @dataclass(frozen=True) @@ -48,28 +39,20 @@ class DatasetVersion: blinding_status: str | None = None -@dataclass(frozen=True) -class DatasetVersionsResult: - """Dataset versions paired with the server's trace id for that call.""" - - items: list[DatasetVersion] - trace_id: str | None = None - - class DatasetFrame: """Lazy dataset reference; fetching requires the originating client to remain open.""" __slots__ = ("_dataset_uuid", "_fetch_data") - def __init__(self, dataset_uuid: str, fetch_data: Callable[[UUID, int | None], FetchDataResult]) -> None: + def __init__(self, dataset_uuid: str, fetch_data: Callable[[UUID, int | None], pd.DataFrame]) -> None: self._dataset_uuid = dataset_uuid self._fetch_data = fetch_data - def head(self, count: int = 6) -> FetchDataResult: + def head(self, count: int = 6) -> pd.DataFrame: """Fetch the first count rows, matching R's default of six rows.""" return self._fetch_data(UUID(self._dataset_uuid), count) - def collect(self) -> FetchDataResult: + def collect(self) -> pd.DataFrame: """Fetch the complete dataset without retaining a previous head limit.""" return self._fetch_data(UUID(self._dataset_uuid), None) @@ -121,7 +104,6 @@ class PaginatedResponse(Generic[T]): # noqa: UP046 total_records: int pagination: Pagination items: list[T] - trace_id: str | None = None @dataclass(frozen=True) @@ -164,7 +146,6 @@ class _PublishEnvelopeResult: checks: ResultChecks = field(default_factory=ResultChecks) errors: list[str] = field(default_factory=list) invalid_records: pd.DataFrame | None = None - trace_id: str | None = None # Flat accessors below are deprecated views onto the envelope, kept so # existing notebooks keep working. Prefer metadata/metrics/checks. @@ -257,8 +238,6 @@ class DatetimeFormatsResult: formats: list[DatetimeFormat] = field(default_factory=list) """The full list of supported formats, in the order returned by the server.""" - trace_id: str | None = None - def all(self) -> list[DatetimeFormat]: """Return every supported format with its type classification.""" return list(self.formats) diff --git a/dataconnect/service/base.py b/dataconnect/service/base.py index eb06a54..4ac7011 100644 --- a/dataconnect/service/base.py +++ b/dataconnect/service/base.py @@ -9,10 +9,9 @@ from dataconnect.models import ( Dataset, - DatasetVersionsResult, + DatasetVersion, DatetimeFormatsResult, DryPublishResult, - FetchDataResult, PaginatedResponse, PublishResult, StudiesResult, @@ -35,14 +34,14 @@ def get_datasets( ) -> PaginatedResponse[Dataset]: ... @abstractmethod - def get_dataset_versions(self, dataset_uuid: UUID) -> DatasetVersionsResult: ... + def get_dataset_versions(self, dataset_uuid: UUID) -> list[DatasetVersion]: ... @abstractmethod def fetch_data( self, dataset_uuid: UUID, first_n_rows: int | None = None, - ) -> FetchDataResult: ... + ) -> pd.DataFrame: ... @abstractmethod def dry_publish( diff --git a/dataconnect/service/default.py b/dataconnect/service/default.py index e23786d..c7bb9cf 100644 --- a/dataconnect/service/default.py +++ b/dataconnect/service/default.py @@ -13,11 +13,10 @@ from dataconnect.models import ( Dataset, DatasetFrame, - DatasetVersionsResult, + DatasetVersion, DatetimeFormat, DatetimeFormatsResult, DryPublishResult, - FetchDataResult, PaginatedResponse, Pagination, PublishResult, @@ -75,6 +74,10 @@ class DefaultDataConnectService(DataConnectService): def __init__(self, transport: Transport) -> None: self._transport = transport + @property + def trace_id(self) -> str | None: + return getattr(self._transport, "trace_id", None) + # DataConnectService def get_studies(self, search_study_name: str | None = None) -> StudiesResult: @@ -94,22 +97,21 @@ def get_studies(self, search_study_name: str | None = None) -> StudiesResult: try: result = self._transport.list_resources(request) - resources = result.resources + resources = result total_records = resources[0].total_records if resources else 0 studies = [resource_to_study(r) for r in resources] - return StudiesResult(total_records=total_records, studies=studies, trace_id=result.trace_id) + return StudiesResult(total_records=total_records, studies=studies) except Exception as ex: raise translate_error(ex) from ex - def get_dataset_versions(self, dataset_uuid: UUID) -> DatasetVersionsResult: + def get_dataset_versions(self, dataset_uuid: UUID) -> list[DatasetVersion]: """List available versions for a dataset. Args: dataset_uuid: UUID of the dataset whose versions are requested. Returns: - A :class:`DatasetVersionsResult` containing the dataset versions, - newest first, and the server's trace id for this call. + The dataset versions, newest first. Raises: ValidationError: If *dataset_uuid* is not a valid UUID (upstream). @@ -122,15 +124,15 @@ def get_dataset_versions(self, dataset_uuid: UUID) -> DatasetVersionsResult: # Return Sorted dataset versions in descending order (newest first) based on the dataset_version field. versions = sorted( - (resource_to_dataset_version(r) for r in result.resources), + (resource_to_dataset_version(r) for r in result), key=lambda dv: dv.dataset_version, reverse=True, ) - return DatasetVersionsResult(items=versions, trace_id=result.trace_id) + return versions except Exception as ex: raise translate_error(ex) from ex - def fetch_data(self, dataset_uuid: UUID, first_n_rows: int | None = None) -> FetchDataResult: + def fetch_data(self, dataset_uuid: UUID, first_n_rows: int | None = None) -> pd.DataFrame: """Fetch data for a dataset""" ticket = DatasetTicket( @@ -140,7 +142,7 @@ def fetch_data(self, dataset_uuid: UUID, first_n_rows: int | None = None) -> Fet try: table = self._transport.get_ticket(ticket) - return FetchDataResult(data=resource_to_fetched_data(table), trace_id=table.trace_id) + return resource_to_fetched_data(table) except TransportError as ex: raise translate_error(ex) from ex @@ -174,7 +176,7 @@ def get_datasets( try: result = self._transport.list_resources(request) - resources = result.resources + resources = result items = [] for resource in resources: dataset = resource_to_dataset(resource) @@ -188,7 +190,6 @@ def get_datasets( total_records=total_records, pagination=Pagination(page=page, page_size=page_size, total_pages=total_pages), items=items, - trace_id=result.trace_id, ) except TransportError as ex: raise translate_error(ex) from ex @@ -362,10 +363,10 @@ def get_datetime_formats( format=fmt, type="datetime" if "HH:mm" in fmt else "date", ) - for fmt in formats_response.formats + for fmt in formats_response ] - return DatetimeFormatsResult(formats=formats, trace_id=formats_response.trace_id) + return DatetimeFormatsResult(formats=formats) def close(self) -> None: """Close the underlying transport connection.""" diff --git a/dataconnect/service/error_handler.py b/dataconnect/service/error_handler.py index ce2c700..86d4d34 100644 --- a/dataconnect/service/error_handler.py +++ b/dataconnect/service/error_handler.py @@ -43,26 +43,54 @@ def translate_error(ex: Exception) -> DataConnectError: if isinstance(ex, TransportAuthenticationError): return AuthenticationError( - error_code=ex.error_code, message=ex.message, timestamp=ex.timestamp, details=error_details + error_code=ex.error_code, + message=ex.message, + timestamp=ex.timestamp, + details=error_details, + trace_id=ex.trace_id, ) if isinstance(ex, TransportAuthorizationError): return AuthorizationError( - error_code=ex.error_code, message=ex.message, timestamp=ex.timestamp, details=error_details + error_code=ex.error_code, + message=ex.message, + timestamp=ex.timestamp, + details=error_details, + trace_id=ex.trace_id, ) if isinstance(ex, TransportValidationError): return ValidationError( - error_code=ex.error_code, message=ex.message, timestamp=ex.timestamp, details=error_details + error_code=ex.error_code, + message=ex.message, + timestamp=ex.timestamp, + details=error_details, + trace_id=ex.trace_id, ) if isinstance(ex, TransportNotFoundError): return NotFoundError( - error_code=ex.error_code, message=ex.message, timestamp=ex.timestamp, details=error_details + error_code=ex.error_code, + message=ex.message, + timestamp=ex.timestamp, + details=error_details, + trace_id=ex.trace_id, ) if isinstance(ex, TransportServerError): - return ServerError(error_code=ex.error_code, message=ex.message, timestamp=ex.timestamp, details=error_details) + return ServerError( + error_code=ex.error_code, + message=ex.message, + timestamp=ex.timestamp, + details=error_details, + trace_id=ex.trace_id, + ) # Non-specific transport error - return DataConnectError(error_code=ex.error_code, message=ex.message, timestamp=ex.timestamp, details=error_details) + return DataConnectError( + error_code=ex.error_code, + message=ex.message, + timestamp=ex.timestamp, + details=error_details, + trace_id=ex.trace_id, + ) diff --git a/dataconnect/service/mappers.py b/dataconnect/service/mappers.py index c5619bb..6d56ef9 100644 --- a/dataconnect/service/mappers.py +++ b/dataconnect/service/mappers.py @@ -147,7 +147,6 @@ def _envelope_to_domain(envelope: PublishEnvelope, result_cls: type[_ResultT]) - ), errors=envelope.errors, invalid_records=envelope.invalid_records, - trace_id=envelope.trace_id, ) diff --git a/dataconnect/transport/arrow_flight/error_handler.py b/dataconnect/transport/arrow_flight/error_handler.py index 328fe40..6595dc7 100644 --- a/dataconnect/transport/arrow_flight/error_handler.py +++ b/dataconnect/transport/arrow_flight/error_handler.py @@ -205,6 +205,7 @@ def parse_dataconnect_error(ex: Exception) -> TransportError: ] error_code = error_data.get("error_code", "SDK_ERROR") + trace_id = error_data.get("trace_id") if error_code.startswith("AUTH_"): return TransportAuthenticationError( @@ -212,6 +213,7 @@ def parse_dataconnect_error(ex: Exception) -> TransportError: message=error_data.get("message") or _UNKNOWN_ERROR, timestamp=error_data.get("timestamp"), details=parsed_details, + trace_id=trace_id, ) if error_code.startswith("AUTHZ_"): @@ -220,6 +222,7 @@ def parse_dataconnect_error(ex: Exception) -> TransportError: message=error_data.get("message") or _UNKNOWN_ERROR, timestamp=error_data.get("timestamp"), details=parsed_details, + trace_id=trace_id, ) if error_code.startswith("VAL_"): @@ -228,6 +231,7 @@ def parse_dataconnect_error(ex: Exception) -> TransportError: message=error_data.get("message") or _UNKNOWN_ERROR, timestamp=error_data.get("timestamp"), details=parsed_details, + trace_id=trace_id, ) if error_code.startswith("RES_"): @@ -236,6 +240,7 @@ def parse_dataconnect_error(ex: Exception) -> TransportError: message=error_data.get("message") or _UNKNOWN_ERROR, timestamp=error_data.get("timestamp"), details=parsed_details, + trace_id=trace_id, ) if error_code.startswith("INT_"): @@ -244,6 +249,7 @@ def parse_dataconnect_error(ex: Exception) -> TransportError: message=error_data.get("message") or _UNKNOWN_ERROR, timestamp=error_data.get("timestamp"), details=parsed_details, + trace_id=trace_id, ) return TransportError( @@ -251,6 +257,7 @@ def parse_dataconnect_error(ex: Exception) -> TransportError: message=error_data.get("message") or _UNKNOWN_ERROR, timestamp=error_data.get("timestamp"), details=parsed_details, + trace_id=trace_id, ) return TransportError(error_code="SDK_ERROR", message=error_message) diff --git a/dataconnect/transport/arrow_flight/transport.py b/dataconnect/transport/arrow_flight/transport.py index 39e495e..c7a75fd 100644 --- a/dataconnect/transport/arrow_flight/transport.py +++ b/dataconnect/transport/arrow_flight/transport.py @@ -12,6 +12,7 @@ import json import platform import subprocess +from collections.abc import Callable from datetime import UTC, datetime from importlib.metadata import version @@ -26,12 +27,10 @@ DatasetTicket, DataTable, DatetimeFormatsRequest, - DatetimeFormatsResponse, DryPublishResponse, PublishRequest, PublishResponse, ResourceInfo, - ResourceListResult, ResourceQuery, ) @@ -51,7 +50,7 @@ def _to_resource_info(info: flight.FlightInfo) -> ResourceInfo: ) -def _to_bytes(table: pa.Table, trace_id: str | None = None) -> DataTable: +def _to_bytes(table: pa.Table) -> DataTable: """Serialize a ``pa.Table`` to a technology-agnostic ``DataTable``. Each record batch is serialized individually as Arrow IPC bytes. @@ -67,7 +66,7 @@ def _to_bytes(table: pa.Table, trace_id: str | None = None) -> DataTable: writer.close() ipc_bytes = sink.getvalue().to_pybytes() - return DataTable(schema_bytes=schema_bytes, ipc_bytes=ipc_bytes, trace_id=trace_id) + return DataTable(schema_bytes=schema_bytes, ipc_bytes=ipc_bytes) def _normalize_arrow_type(dtype: pa.DataType) -> pa.DataType: @@ -89,6 +88,68 @@ def _normalize_arrow_type(dtype: pa.DataType) -> pa.DataType: return dtype +_TRACE_ID_RESPONSE_HEADER = "x-dataconnect-trace-id" + + +class _TraceClientMiddleware(flight.ClientMiddleware): + def __init__( + self, + capture_trace_id: Callable[[str | None], None], + capture_payload_trace_id: Callable[[object], None], + ) -> None: + self._capture_trace_id = capture_trace_id + self._capture_payload_trace_id = capture_payload_trace_id + + def received_headers(self, headers: dict[str, list[str | bytes]]) -> None: + values = headers.get(_TRACE_ID_RESPONSE_HEADER) + if not values: + return + + trace_id = values[0] + if isinstance(trace_id, bytes): + trace_id = trace_id.decode("utf-8", errors="replace") + if trace_id: + self._capture_trace_id(trace_id) + + def call_completed(self, exception: Exception | None) -> None: + if exception is None: + return + + message = str(exception) + separator = message.find("::") + if separator < 0: + return + + payload_text = message[separator + 2 :] + payload_start = payload_text.find("{") + if payload_start < 0: + return + + try: + payload, _ = json.JSONDecoder().raw_decode(payload_text[payload_start:]) + except json.JSONDecodeError: + return + + if isinstance(payload, dict): + self._capture_payload_trace_id(payload.get("trace_id")) + + +class _TraceClientMiddlewareFactory(flight.ClientMiddlewareFactory): + def __init__( + self, + begin_call: Callable[[], None], + capture_trace_id: Callable[[str | None], None], + capture_payload_trace_id: Callable[[object], None], + ) -> None: + self._begin_call = begin_call + self._capture_trace_id = capture_trace_id + self._capture_payload_trace_id = capture_payload_trace_id + + def start_call(self, _info: object) -> _TraceClientMiddleware: + self._begin_call() + return _TraceClientMiddleware(self._capture_trace_id, self._capture_payload_trace_id) + + # Maps service-layer action names to the flight_type value the Arrow Flight server expects. _ACTION_FLIGHT_TYPE: dict[str, str] = { "studies.list": "STUDIES", @@ -118,6 +179,10 @@ def __init__( token: Optional Bearer token appended to every request header. """ self._call_headers: list[tuple[bytes, bytes]] = [] + self._trace_id: str | None = None + self._trace_middleware = _TraceClientMiddlewareFactory( + self._begin_call, self._capture_trace_id, self._capture_payload_trace_id + ) scheme = "grpc+tls" if use_tls else "grpc" uri = f"{scheme}://{host}:{port}" @@ -174,9 +239,9 @@ def _get_client(self, uri: str, use_tls: bool) -> flight.FlightClient: pem_parts.append("\n".join(lines)) pem_certs = "\n".join(pem_parts).encode("utf-8") - client = flight.FlightClient(uri, tls_root_certs=pem_certs) + client = flight.FlightClient(uri, tls_root_certs=pem_certs, middleware=[self._trace_middleware]) else: - client = flight.FlightClient(uri) + client = flight.FlightClient(uri, middleware=[self._trace_middleware]) return client @@ -184,9 +249,24 @@ def _options(self) -> flight.FlightCallOptions: """Return call options containing the configured request headers.""" return flight.FlightCallOptions(headers=self._call_headers) + @property + def trace_id(self) -> str | None: + return self._trace_id + + def _begin_call(self) -> None: + self._trace_id = None + + def _capture_trace_id(self, trace_id: str | None) -> None: + if trace_id: + self._trace_id = trace_id + + def _capture_payload_trace_id(self, trace_id: object) -> None: + if self._trace_id is None and isinstance(trace_id, str) and trace_id: + self._trace_id = trace_id + # Transport - def list_resources(self, request: ResourceQuery) -> ResourceListResult: + def list_resources(self, request: ResourceQuery) -> list[ResourceInfo]: """Translate the action name to Arrow Flight criteria and return resource records.""" flight_type = _ACTION_FLIGHT_TYPE.get(request.action) @@ -204,14 +284,15 @@ def list_resources(self, request: ResourceQuery) -> ResourceListResult: try: raw_flights = list(self._client.list_flights(criteria, self._options())) - trace_id = None if raw_flights and raw_flights[0].app_metadata: try: - trace_id = json.loads(raw_flights[0].app_metadata.decode("utf-8")).get("trace_id") + metadata = json.loads(raw_flights[0].app_metadata.decode("utf-8")) + if isinstance(metadata, dict): + self._capture_payload_trace_id(metadata.get("trace_id")) except (json.JSONDecodeError, UnicodeDecodeError): - trace_id = None + pass - return ResourceListResult(resources=[_to_resource_info(f) for f in raw_flights], trace_id=trace_id) + return [_to_resource_info(f) for f in raw_flights] except Exception as ex: raise parse_dataconnect_error(ex) from ex @@ -234,13 +315,12 @@ def get_ticket(self, ticket: DatasetTicket) -> DataTable: except flight.FlightError as ex: raise parse_dataconnect_error(ex) from ex - trace_id = None if table.schema.metadata: raw_trace_id = table.schema.metadata.get(b"trace_id") if raw_trace_id is not None: - trace_id = raw_trace_id.decode("utf-8") + self._capture_payload_trace_id(raw_trace_id.decode("utf-8")) - return _to_bytes(pa.Table.from_batches(batches, schema=table.schema), trace_id=trace_id) + return _to_bytes(pa.Table.from_batches(batches, schema=table.schema)) except Exception as ex: raise parse_dataconnect_error(ex) from ex @@ -371,7 +451,7 @@ def publish_dataset(self, publish_request: PublishRequest) -> PublishResponse: except Exception as ex: raise parse_dataconnect_error(ex) from ex - def get_datetime_formats(self, request: DatetimeFormatsRequest) -> DatetimeFormatsResponse: + def get_datetime_formats(self, request: DatetimeFormatsRequest) -> list[str]: """Invoke the Arrow Flight ``get_datetime_formats`` action and return the format list. The server expects a JSON body of the form @@ -398,10 +478,15 @@ def get_datetime_formats(self, request: DatetimeFormatsRequest) -> DatetimeForma results = list(self._client.do_action(action, self._options())) if not results: - return DatetimeFormatsResponse(formats=[], trace_id=None) + return [] response = json.loads(results[0].body.to_pybytes().decode("utf-8")) - return DatetimeFormatsResponse(formats=response.get("formats", []), trace_id=response.get("trace_id")) + if isinstance(response, list): + return response + if isinstance(response, dict): + self._capture_payload_trace_id(response.get("trace_id")) + return response.get("formats", []) + return [] except Exception as ex: raise parse_dataconnect_error(ex) from ex diff --git a/dataconnect/transport/base.py b/dataconnect/transport/base.py index 956c828..1ce7078 100644 --- a/dataconnect/transport/base.py +++ b/dataconnect/transport/base.py @@ -13,11 +13,10 @@ DatasetTicket, DataTable, DatetimeFormatsRequest, - DatetimeFormatsResponse, DryPublishResponse, PublishRequest, PublishResponse, - ResourceListResult, + ResourceInfo, ResourceQuery, ) @@ -26,15 +25,14 @@ class Transport(ABC): """Minimal abstract transport for DataConnect operations.""" @abstractmethod - def list_resources(self, request: ResourceQuery) -> ResourceListResult: + def list_resources(self, request: ResourceQuery) -> list[ResourceInfo]: """List available data resources matching the given query. The transport does not interpret the action name or body — that is the service layer's responsibility. Returns: - A :class:`ResourceListResult` with the matched resources and the - server's trace id for this call, or ``None`` if not provided. + The matched resources. """ @abstractmethod @@ -62,7 +60,7 @@ def publish_dataset(self, publish_request: PublishRequest) -> PublishResponse: """ @abstractmethod - def get_datetime_formats(self, request: DatetimeFormatsRequest) -> DatetimeFormatsResponse: + def get_datetime_formats(self, request: DatetimeFormatsRequest) -> list[str]: """Return the supported datetime format strings for the project. The transport does not interpret ``format_type`` — that is the service @@ -70,8 +68,7 @@ def get_datetime_formats(self, request: DatetimeFormatsRequest) -> DatetimeForma already-filtered list of format strings. Returns: - A :class:`DatetimeFormatsResponse` with the filtered formats and - the server's trace id for this call, or ``None`` if not provided. + The filtered datetime formats. """ @abstractmethod diff --git a/dataconnect/transport/errors.py b/dataconnect/transport/errors.py index 3fd9578..f9d6b5d 100644 --- a/dataconnect/transport/errors.py +++ b/dataconnect/transport/errors.py @@ -46,6 +46,7 @@ class TransportError(Exception): message: str timestamp: str | None = None details: list[ErrorDetail] | None = None + trace_id: str | None = None class TransportAuthenticationError(TransportError): diff --git a/dataconnect/transport/models.py b/dataconnect/transport/models.py index a389a83..5befc2d 100644 --- a/dataconnect/transport/models.py +++ b/dataconnect/transport/models.py @@ -53,26 +53,16 @@ class ResourceInfo: schema_bytes: bytes -@dataclass(frozen=True) -class ResourceListResult: - """Result of a resource-listing call (studies/datasets/dataset_versions).""" - - resources: list[ResourceInfo] - trace_id: str | None = None - - @dataclass(frozen=True) class DataTable: """Technology-agnostic representation of a fetched data result. ``schema_bytes`` holds the Arrow IPC-serialized schema. ``ipc_bytes`` holds the full Arrow IPC stream (schema + all batches). - ``trace_id`` is the server's trace id for the call that produced this result, if any. """ schema_bytes: bytes ipc_bytes: bytes - trace_id: str | None = None @dataclass(frozen=True) @@ -90,14 +80,6 @@ class DatetimeFormatsRequest: format_type: str = "all" -@dataclass(frozen=True) -class DatetimeFormatsResponse: - """Result of a get_datetime_formats call.""" - - formats: list[str] - trace_id: str | None = None - - @dataclass(frozen=True) class PublishRequest: """A publish request containing the input configuration and the dataset to be published. @@ -161,8 +143,6 @@ class PublishEnvelope: errors: list[str] = field(default_factory=list) invalid_records: pd.DataFrame | None = None """Populated from the Arrow IPC channel, not from the JSON payload.""" - trace_id: str | None = None - """Populated from the JSON payload.""" @classmethod def from_json(cls, payload: dict, invalid_records: pd.DataFrame | None = None) -> PublishEnvelope: @@ -198,7 +178,6 @@ def from_json(cls, payload: dict, invalid_records: pd.DataFrame | None = None) - ), errors=payload.get("errors") or [], invalid_records=invalid_records, - trace_id=payload.get("trace_id"), ) diff --git a/tests/test_dry_publish.py b/tests/test_dry_publish.py index f8210c3..3bcc74a 100644 --- a/tests/test_dry_publish.py +++ b/tests/test_dry_publish.py @@ -18,7 +18,7 @@ import pytest from dataconnect.exceptions import ValidationError -from dataconnect.models import DatetimeFormatsResult, DryPublishResult +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 ( @@ -30,11 +30,11 @@ from dataconnect.transport.models import ( DatasetTicket, DataTable, - DatetimeFormatsResponse, + DatetimeFormatsRequest, DryPublishResponse, PublishRequest, PublishResponse, - ResourceListResult, + ResourceInfo, ResourceQuery, ResponseChecks, ResponseMetadata, @@ -85,7 +85,6 @@ def _make_dry_publish_response(**overrides: object) -> DryPublishResponse: ), errors=flat["errors"], invalid_records=flat["invalid_records"], - trace_id=overrides.get("trace_id"), # type: ignore[arg-type] ) @@ -227,13 +226,9 @@ def test_returns_dry_publish_result_instance(self) -> None: result = dry_publish_response_to_domain(_make_dry_publish_response()) assert isinstance(result, DryPublishResult) - def test_trace_id_mapped(self) -> None: - result = dry_publish_response_to_domain(_make_dry_publish_response(trace_id="trace-dry-1")) - assert result.trace_id == "trace-dry-1" - - def test_trace_id_defaults_to_none(self) -> None: + def test_trace_id_is_not_added_to_result_contract(self) -> None: result = dry_publish_response_to_domain(_make_dry_publish_response()) - assert result.trace_id is None + assert not hasattr(result, "trace_id") # --------------------------------------------------------------------------- @@ -253,8 +248,8 @@ def __init__( self._raise = raise_error self.last_request: PublishRequest | None = None - def list_resources(self, request: ResourceQuery) -> ResourceListResult: - return ResourceListResult(resources=[], trace_id=None) + def list_resources(self, request: ResourceQuery) -> list[ResourceInfo]: + return [] def get_ticket(self, ticket: DatasetTicket) -> DataTable: raise NotImplementedError @@ -268,7 +263,7 @@ def dry_publish_dataset(self, publish_request: PublishRequest) -> DryPublishResp def publish_dataset(self, publish_request: PublishRequest) -> PublishResponse: raise NotImplementedError - def get_datetime_formats(self, request: DatetimeFormatsResult) -> DatetimeFormatsResponse: # type: ignore[override] + def get_datetime_formats(self, request: DatetimeFormatsRequest) -> list[str]: raise NotImplementedError def close(self) -> None: @@ -302,11 +297,6 @@ def test_returns_dry_publish_result_instance(self) -> None: result = service.dry_publish(**_default_dry_publish_args()) assert isinstance(result, DryPublishResult) - def test_trace_id_is_propagated_from_transport(self) -> None: - service, _ = _make_service(dry_publish_return=_make_dry_publish_response(trace_id="trace-dry-svc-1")) - result = service.dry_publish(**_default_dry_publish_args()) - assert result.trace_id == "trace-dry-svc-1" - 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 @@ -529,13 +519,6 @@ def test_returns_dry_publish_response_instance(self) -> None: result = transport.dry_publish_dataset(PublishRequest(input_config="{}", data=pd.DataFrame({"x": [1]}))) assert isinstance(result, DryPublishResponse) - def test_trace_id_is_read_from_json_payload(self) -> None: - transport = _make_flight_transport() - _wire_do_put(transport, {**_VALID_JSON_RESP, "trace_id": "trace-put-1"}) - - result = transport.dry_publish_dataset(PublishRequest(input_config="{}", data=pd.DataFrame({"x": [1]}))) - assert result.trace_id == "trace-put-1" - def test_status_parsed_from_json(self) -> None: transport = _make_flight_transport() _wire_do_put(transport, {**_VALID_JSON_RESP, "success": False}) diff --git a/tests/test_fetch_data.py b/tests/test_fetch_data.py index f1d5a1b..23410ff 100644 --- a/tests/test_fetch_data.py +++ b/tests/test_fetch_data.py @@ -22,7 +22,7 @@ TransportNotFoundError, TransportServerError, ) -from dataconnect.transport.models import DatasetTicket, DataTable, ResourceListResult, ResourceQuery +from dataconnect.transport.models import DatasetTicket, DataTable, ResourceQuery # --------------------------------------------------------------------------- # Fake transport @@ -41,8 +41,8 @@ def __init__( self._get_ticket_error = get_ticket_error self.last_ticket: DatasetTicket | None = None - def list_resources(self, request: ResourceQuery) -> ResourceListResult: - return ResourceListResult(resources=[], trace_id=None) + def list_resources(self, request: ResourceQuery) -> list: + return [] def get_ticket(self, ticket: DatasetTicket) -> DataTable: self.last_ticket = ticket @@ -55,7 +55,7 @@ def close(self) -> None: return None -def _make_ipc_table(data: dict, trace_id: str | None = None) -> DataTable: +def _make_ipc_table(data: dict) -> DataTable: arrow_table = pa.table(data) sink = pa.BufferOutputStream() writer = pa.ipc.new_stream(sink, arrow_table.schema) @@ -64,7 +64,6 @@ def _make_ipc_table(data: dict, trace_id: str | None = None) -> DataTable: return DataTable( schema_bytes=arrow_table.schema.serialize().to_pybytes(), ipc_bytes=sink.getvalue().to_pybytes(), - trace_id=trace_id, ) @@ -81,9 +80,9 @@ def test_fetch_data_returns_dataframe_with_correct_values() -> None: result = service.fetch_data(dataset_uuid) - assert isinstance(result.data, pd.DataFrame) - assert result.data["subject_id"].tolist() == source["subject_id"] - assert result.data["age"].tolist() == source["age"] + 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: @@ -98,26 +97,6 @@ def test_fetch_data_builds_correct_ticket() -> None: assert transport.last_ticket.limit == 10 -def test_fetch_data_returns_trace_id_from_transport() -> None: - dataset_uuid = UUID("073410b6-79be-3e7d-ae37-92f6e054013e") - transport = _FakeTransport(data_table=_make_ipc_table({"x": [1]}, trace_id="abc123")) - service = DefaultDataConnectService(transport) - - result = service.fetch_data(dataset_uuid) - - assert result.trace_id == "abc123" - - -def test_fetch_data_trace_id_defaults_to_none() -> None: - dataset_uuid = UUID("073410b6-79be-3e7d-ae37-92f6e054013e") - transport = _FakeTransport(data_table=_make_ipc_table({"x": [1]})) - service = DefaultDataConnectService(transport) - - result = service.fetch_data(dataset_uuid) - - assert result.trace_id is None - - 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]})) @@ -168,9 +147,9 @@ def test_fetch_data_returns_empty_dataframe_for_empty_table() -> None: result = service.fetch_data(dataset_uuid) - assert isinstance(result.data, pd.DataFrame) - assert len(result.data) == 0 - assert "col" in result.data.columns + assert isinstance(result, pd.DataFrame) + assert len(result) == 0 + assert "col" in result.columns def test_fetch_data_translates_authentication_error() -> None: diff --git a/tests/test_get_datasets_paginated.py b/tests/test_get_datasets_paginated.py index bad16a8..9af6d24 100644 --- a/tests/test_get_datasets_paginated.py +++ b/tests/test_get_datasets_paginated.py @@ -21,7 +21,6 @@ DatasetTicket, DataTable, ResourceInfo, - ResourceListResult, ResourceQuery, ) @@ -35,14 +34,14 @@ def __init__( ) -> None: self._resources = resources or [] self._error = error - self._trace_id = trace_id + self.trace_id = trace_id self.last_request: ResourceQuery | None = None - def list_resources(self, request: ResourceQuery) -> ResourceListResult: + def list_resources(self, request: ResourceQuery) -> list[ResourceInfo]: self.last_request = request if self._error is not None: raise self._error - return ResourceListResult(resources=self._resources, trace_id=self._trace_id) + return self._resources def do_get(self, request: ResourceQuery) -> None: pass @@ -169,17 +168,17 @@ def test_returns_trace_id_from_transport(self) -> None: transport = _FakeTransport(resources=[], trace_id="trace-datasets-1") service = DefaultDataConnectService(transport) - result = service.get_datasets(study_environment_uuid=_STUDY_ENV_UUID) + service.get_datasets(study_environment_uuid=_STUDY_ENV_UUID) - assert result.trace_id == "trace-datasets-1" + assert service.trace_id == "trace-datasets-1" def test_trace_id_defaults_to_none(self) -> None: transport = _FakeTransport(resources=[]) service = DefaultDataConnectService(transport) - result = service.get_datasets(study_environment_uuid=_STUDY_ENV_UUID) + service.get_datasets(study_environment_uuid=_STUDY_ENV_UUID) - assert result.trace_id is None + assert service.trace_id is None def test_multiple_items_returned(self) -> None: resources = [ @@ -252,12 +251,12 @@ def __init__(self, pages: dict[int, list[dict[str, object]]]) -> None: self.lock = Lock() self.fetch_error: TransportError | None = None - def list_resources(self, request: ResourceQuery) -> ResourceListResult: + def list_resources(self, request: ResourceQuery) -> list[ResourceInfo]: self.requests.append(request) body = json.loads(request.body) total_records = sum(len(payloads) for payloads in self.pages.values()) resources = [_dataset_resource(payload, total_records) for payload in self.pages[body["page"]]] - return ResourceListResult(resources=resources, trace_id=None) + return resources def get_ticket(self, ticket: DatasetTicket) -> DataTable: if self.closed: @@ -357,13 +356,13 @@ def test_frame_fetches_only_on_demand_and_binds_each_dataset() -> None: assert second.frame is not None assert first.frame is not second.frame - preview = first.frame.head(3).data + preview = first.frame.head(3) assert isinstance(preview, pd.DataFrame) assert preview["value"].tolist() == [0, 1, 2] assert preview["dataset_uuid"].tolist() == [_DATASET_UUID] * 3 - assert len(first.frame.head().data) == 6 - assert len(first.frame.collect().data) == 8 - assert second.frame.head(1).data["dataset_uuid"].tolist() == [_OTHER_DATASET_UUID] + assert len(first.frame.head()) == 6 + assert len(first.frame.collect()) == 8 + assert second.frame.head(1)["dataset_uuid"].tolist() == [_OTHER_DATASET_UUID] assert transport.tickets == [ DatasetTicket(dataset_uuid=_DATASET_UUID, limit=3), DatasetTicket(dataset_uuid=_DATASET_UUID, limit=6), @@ -387,7 +386,7 @@ def test_frame_copy_and_inspection_do_not_copy_or_fetch_connection() -> None: for name, value in _METADATA.items(): assert encoded[name] == value assert transport.tickets == [] - assert len(encoded["frame"].head(1).data) == 1 + assert len(encoded["frame"].head(1)) == 1 def test_dataset_constructor_and_equality_remain_independent_of_frame() -> None: @@ -494,12 +493,12 @@ def test_dataset_versions_response_is_unchanged_before_and_after_listing() -> No ] listed = client.get_dataset_versions(UUID(_DATASET_UUID)) - assert [asdict(item) for item in listed.items] == expected - assert listed.items[0].blinding_status is None + assert [asdict(item) for item in listed] == expected + assert listed[0].blinding_status is None transport._resources = [_dataset_resource({**_IDENTIFIERS, **_METADATA})] assert client.get_datasets(_STUDY_ENV_UUID).items[0].frame is not None transport._resources = versions - assert [asdict(item) for item in client.get_dataset_versions(UUID(_DATASET_UUID)).items] == expected + assert [asdict(item) for item in client.get_dataset_versions(UUID(_DATASET_UUID))] == expected assert transport.last_request is not None assert transport.last_request.action == "dataset_versions.list" assert json.loads(transport.last_request.body) == {"dataset_uuid": _DATASET_UUID} @@ -512,6 +511,6 @@ def test_get_dataset_versions_returns_trace_id_from_transport() -> None: ) service = DefaultDataConnectService(transport) - result = service.get_dataset_versions(UUID(_DATASET_UUID)) + service.get_dataset_versions(UUID(_DATASET_UUID)) - assert result.trace_id == "trace-versions-1" + assert service.trace_id == "trace-versions-1" diff --git a/tests/test_get_datetime_formats.py b/tests/test_get_datetime_formats.py index 74df91a..7bfbd49 100644 --- a/tests/test_get_datetime_formats.py +++ b/tests/test_get_datetime_formats.py @@ -25,11 +25,9 @@ DatasetTicket, DataTable, DatetimeFormatsRequest, - DatetimeFormatsResponse, DryPublishResponse, PublishRequest, PublishResponse, - ResourceListResult, ResourceQuery, ) @@ -57,11 +55,11 @@ def __init__( ) -> None: self._formats = formats if formats is not None else list(_SAMPLE_FORMATS) self._raise = raise_error - self._trace_id = trace_id + self.trace_id = trace_id self.last_request: DatetimeFormatsRequest | None = None - def list_resources(self, request: ResourceQuery) -> ResourceListResult: - return ResourceListResult(resources=[], trace_id=None) + def list_resources(self, request: ResourceQuery) -> list: + return [] def get_ticket(self, ticket: DatasetTicket) -> DataTable: raise NotImplementedError @@ -72,11 +70,11 @@ def dry_publish_dataset(self, publish_request: PublishRequest) -> DryPublishResp def publish_dataset(self, publish_request: PublishRequest) -> PublishResponse: raise NotImplementedError - def get_datetime_formats(self, request: DatetimeFormatsRequest) -> DatetimeFormatsResponse: + def get_datetime_formats(self, request: DatetimeFormatsRequest) -> list[str]: self.last_request = request if self._raise is not None: raise self._raise - return DatetimeFormatsResponse(formats=self._formats, trace_id=self._trace_id) + return self._formats def close(self) -> None: pass @@ -255,13 +253,13 @@ def test_empty_server_response_yields_empty_result(self) -> None: def test_trace_id_is_propagated_from_transport(self) -> None: service, _ = _make_service(trace_id="trace-fmt-1") - result = service.get_datetime_formats(project_token="tok") - assert result.trace_id == "trace-fmt-1" + service.get_datetime_formats(project_token="tok") + assert service.trace_id == "trace-fmt-1" def test_trace_id_defaults_to_none(self) -> None: service, _ = _make_service() - result = service.get_datetime_formats(project_token="tok") - assert result.trace_id is None + service.get_datetime_formats(project_token="tok") + assert service.trace_id is None # --- error translation --- @@ -288,9 +286,9 @@ def _make_flight_transport() -> ArrowFlightTransport: return ArrowFlightTransport(host="localhost", port=5005, use_tls=False) -def _make_action_result(formats: list[str], trace_id: str | None = None) -> MagicMock: +def _make_action_result(formats: list[str]) -> MagicMock: body = MagicMock() - body.to_pybytes.return_value = json.dumps({"formats": formats, "trace_id": trace_id}).encode("utf-8") + body.to_pybytes.return_value = json.dumps(formats).encode("utf-8") result = MagicMock() result.body = body return result @@ -326,7 +324,7 @@ def test_returns_decoded_format_list(self) -> None: response = transport.get_datetime_formats(DatetimeFormatsRequest(project_token="tok", format_type="all")) - assert response.formats == _SAMPLE_FORMATS + assert response == _SAMPLE_FORMATS def test_empty_result_iterator_returns_empty_list(self) -> None: transport = _make_flight_transport() @@ -334,15 +332,7 @@ def test_empty_result_iterator_returns_empty_list(self) -> None: response = transport.get_datetime_formats(DatetimeFormatsRequest(project_token="tok", format_type="all")) - assert response.formats == [] - - def test_trace_id_is_read_from_json_payload(self) -> None: - transport = _make_flight_transport() - transport._client.do_action.return_value = iter([_make_action_result(_SAMPLE_FORMATS, trace_id="trace-fmt-2")]) - - response = transport.get_datetime_formats(DatetimeFormatsRequest(project_token="tok", format_type="all")) - - assert response.trace_id == "trace-fmt-2" + assert response == [] def test_underlying_exception_is_translated_to_transport_error(self) -> None: from dataconnect.transport.errors import TransportError diff --git a/tests/test_publish.py b/tests/test_publish.py index 2f8ea10..85700bf 100644 --- a/tests/test_publish.py +++ b/tests/test_publish.py @@ -30,11 +30,10 @@ DatasetTicket, DataTable, DatetimeFormatsRequest, - DatetimeFormatsResponse, DryPublishResponse, PublishRequest, PublishResponse, - ResourceListResult, + ResourceInfo, ResourceQuery, ResponseMetadata, ResponseMetrics, @@ -73,7 +72,6 @@ def _make_publish_response(**overrides: object) -> PublishResponse: total_duplicate_rows=flat["duplicate_record_count"], ), invalid_records=flat["invalid_records"], - trace_id=overrides.get("trace_id"), # type: ignore[arg-type] ) @@ -148,13 +146,9 @@ def test_returns_publish_result_instance(self) -> None: result = publish_response_to_domain(_make_publish_response()) assert isinstance(result, PublishResult) - def test_trace_id_mapped(self) -> None: - result = publish_response_to_domain(_make_publish_response(trace_id="trace-pub-1")) - assert result.trace_id == "trace-pub-1" - - def test_trace_id_defaults_to_none(self) -> None: + def test_trace_id_is_not_added_to_result_contract(self) -> None: result = publish_response_to_domain(_make_publish_response()) - assert result.trace_id is None + assert not hasattr(result, "trace_id") # --------------------------------------------------------------------------- @@ -174,8 +168,8 @@ def __init__( self._raise = raise_error self.last_request: PublishRequest | None = None - def list_resources(self, request: ResourceQuery) -> ResourceListResult: - return ResourceListResult(resources=[], trace_id=None) + def list_resources(self, request: ResourceQuery) -> list[ResourceInfo]: + return [] def get_ticket(self, ticket: DatasetTicket) -> DataTable: raise NotImplementedError @@ -189,7 +183,7 @@ def publish_dataset(self, publish_request: PublishRequest) -> PublishResponse: raise self._raise return self._return # type: ignore[return-value] - def get_datetime_formats(self, request: DatetimeFormatsRequest) -> DatetimeFormatsResponse: # type: ignore[override] + def get_datetime_formats(self, request: DatetimeFormatsRequest) -> list[str]: raise NotImplementedError def close(self) -> None: @@ -223,11 +217,6 @@ def test_returns_publish_result_instance(self) -> None: result = service.publish(**_default_publish_args()) assert isinstance(result, PublishResult) - def test_trace_id_is_propagated_from_transport(self) -> None: - service, _ = _make_service(publish_return=_make_publish_response(trace_id="trace-pub-svc-1")) - result = service.publish(**_default_publish_args()) - assert result.trace_id == "trace-pub-svc-1" - def test_successful_status_is_propagated(self) -> None: service, _ = _make_service(publish_return=_make_publish_response(status=True)) assert service.publish(**_default_publish_args()).status is True @@ -450,13 +439,6 @@ def test_returns_publish_response_instance(self) -> None: result = transport.publish_dataset(PublishRequest(input_config="{}", data=pd.DataFrame({"x": [1]}))) assert isinstance(result, PublishResponse) - def test_trace_id_is_read_from_json_payload(self) -> None: - transport = _make_flight_transport() - _wire_do_put(transport, {**_VALID_JSON_RESP, "trace_id": "trace-put-2"}) - - result = transport.publish_dataset(PublishRequest(input_config="{}", data=pd.DataFrame({"x": [1]}))) - assert result.trace_id == "trace-put-2" - def test_status_parsed_from_json(self) -> None: transport = _make_flight_transport() _wire_do_put(transport, {**_VALID_JSON_RESP, "success": False}) diff --git a/tests/test_transport_trace_metadata.py b/tests/test_transport_trace_metadata.py index b7d9e02..bb658d4 100644 --- a/tests/test_transport_trace_metadata.py +++ b/tests/test_transport_trace_metadata.py @@ -5,6 +5,8 @@ import pyarrow as pa import pytest +from dataconnect.service.error_handler import translate_error +from dataconnect.transport.arrow_flight.error_handler import parse_dataconnect_error from dataconnect.transport.arrow_flight.transport import ArrowFlightTransport from dataconnect.transport.models import DatasetTicket, ResourceQuery @@ -30,8 +32,8 @@ def test_list_resources_reads_trace_id_from_app_metadata() -> None: result = transport.list_resources(ResourceQuery(action="studies.list")) - assert result.trace_id == "trace-list-1" - assert len(result.resources) == 1 + assert transport.trace_id == "trace-list-1" + assert len(result) == 1 @pytest.mark.parametrize("app_metadata", [None, b"", b"{}", b"not-json", b"\xff"]) @@ -39,9 +41,11 @@ def test_list_resources_returns_no_trace_id_for_missing_or_malformed_metadata(ap transport = _make_transport() transport._client.list_flights.return_value = [_make_flight_info(app_metadata)] + transport._begin_call() result = transport.list_resources(ResourceQuery(action="studies.list")) - assert result.trace_id is None + assert transport.trace_id is None + assert isinstance(result, list) def test_get_ticket_reads_trace_id_from_schema_metadata() -> None: @@ -53,7 +57,9 @@ def test_get_ticket_reads_trace_id_from_schema_metadata() -> None: result = transport.get_ticket(DatasetTicket(dataset_uuid="dataset-1")) - assert result.trace_id == "trace-get-1" + assert transport.trace_id == "trace-get-1" + assert isinstance(result.schema_bytes, bytes) + assert isinstance(result.ipc_bytes, bytes) def test_get_ticket_returns_no_trace_id_when_schema_metadata_is_absent() -> None: @@ -65,4 +71,46 @@ def test_get_ticket_returns_no_trace_id_when_schema_metadata_is_absent() -> None result = transport.get_ticket(DatasetTicket(dataset_uuid="dataset-1")) - assert result.trace_id is None + assert transport.trace_id is None + assert isinstance(result.schema_bytes, bytes) + + +def test_client_middleware_captures_response_trace_header_and_clears_previous_value() -> None: + transport = _make_transport() + transport._trace_id = "previous-trace" + + middleware = transport._trace_middleware.start_call(None) + + assert transport.trace_id is None + middleware.received_headers({"x-dataconnect-trace-id": ["current-trace"]}) + assert transport.trace_id == "current-trace" + + +def test_client_middleware_reads_trace_id_from_structured_pre_handler_error() -> None: + transport = _make_transport() + middleware = transport._trace_middleware.start_call(None) + + middleware.call_completed( + RuntimeError('AUTH_001::{"error_code":"AUTH_001","trace_id":"structured-trace"}. Detail: Unauthenticated') + ) + + assert transport.trace_id == "structured-trace" + + +def test_structured_error_trace_id_is_preserved_and_printed() -> None: + exception = RuntimeError( + 'RES_002::{"error_code":"RES_002","message":"Dataset lookup failed",' + '"details":[{"field":"dataset_uuid",' + '"expected":"Review and provide the correct dataset_uuid that is associated with a valid Study Environment."}],' + '"trace_id":"structured-error-trace"}' + ) + + transport_error = parse_dataconnect_error(exception) + public_error = translate_error(transport_error) + + assert transport_error.trace_id == "structured-error-trace" + assert public_error.trace_id == "structured-error-trace" + assert str(public_error).endswith( + " Expected: Review and provide the correct dataset_uuid that is associated with a valid Study Environment.\n" + " Trace ID: structured-error-trace" + ) From f57d84a188df2ac48d0024084ca1e695962993ed Mon Sep 17 00:00:00 2001 From: Gaurav Londhe Date: Wed, 30 Sep 2026 16:40:41 +0530 Subject: [PATCH 3/3] feat: resolved copilot comments --- dataconnect/transport/arrow_flight/transport.py | 9 ++++++++- tests/test_dry_publish.py | 9 +++++++++ tests/test_publish.py | 9 +++++++++ tests/test_transport_trace_metadata.py | 13 +++++++++++++ 4 files changed, 39 insertions(+), 1 deletion(-) diff --git a/dataconnect/transport/arrow_flight/transport.py b/dataconnect/transport/arrow_flight/transport.py index c7a75fd..f76b2d2 100644 --- a/dataconnect/transport/arrow_flight/transport.py +++ b/dataconnect/transport/arrow_flight/transport.py @@ -318,7 +318,12 @@ def get_ticket(self, ticket: DatasetTicket) -> DataTable: if table.schema.metadata: raw_trace_id = table.schema.metadata.get(b"trace_id") if raw_trace_id is not None: - self._capture_payload_trace_id(raw_trace_id.decode("utf-8")) + try: + trace_id = raw_trace_id.decode("utf-8") + except UnicodeDecodeError: + pass + else: + self._capture_payload_trace_id(trace_id) return _to_bytes(pa.Table.from_batches(batches, schema=table.schema)) @@ -379,6 +384,7 @@ def dry_publish_dataset(self, publish_request: PublishRequest) -> DryPublishResp # 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()) + self._capture_payload_trace_id(json_result.get("trace_id")) # Read the Arrow table returned by the server upon successful publishing metadata_buf = reader.read() @@ -436,6 +442,7 @@ def publish_dataset(self, publish_request: PublishRequest) -> PublishResponse: # 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()) + self._capture_payload_trace_id(json_result.get("trace_id")) # Read the Arrow table returned by the server upon successful publishing metadata_buf = reader.read() diff --git a/tests/test_dry_publish.py b/tests/test_dry_publish.py index 3bcc74a..86ca44d 100644 --- a/tests/test_dry_publish.py +++ b/tests/test_dry_publish.py @@ -519,6 +519,15 @@ def test_returns_dry_publish_response_instance(self) -> None: result = transport.dry_publish_dataset(PublishRequest(input_config="{}", data=pd.DataFrame({"x": [1]}))) assert isinstance(result, DryPublishResponse) + def test_captures_trace_id_from_response_payload_without_changing_result(self) -> None: + transport = _make_flight_transport() + _wire_do_put(transport, {**_VALID_JSON_RESP, "trace_id": "dry-publish-trace"}) + + result = transport.dry_publish_dataset(PublishRequest(input_config="{}", data=pd.DataFrame({"x": [1]}))) + + assert transport.trace_id == "dry-publish-trace" + assert not hasattr(result, "trace_id") + def test_status_parsed_from_json(self) -> None: transport = _make_flight_transport() _wire_do_put(transport, {**_VALID_JSON_RESP, "success": False}) diff --git a/tests/test_publish.py b/tests/test_publish.py index 85700bf..4a9b342 100644 --- a/tests/test_publish.py +++ b/tests/test_publish.py @@ -439,6 +439,15 @@ def test_returns_publish_response_instance(self) -> None: result = transport.publish_dataset(PublishRequest(input_config="{}", data=pd.DataFrame({"x": [1]}))) assert isinstance(result, PublishResponse) + def test_captures_trace_id_from_response_payload_without_changing_result(self) -> None: + transport = _make_flight_transport() + _wire_do_put(transport, {**_VALID_JSON_RESP, "trace_id": "publish-trace"}) + + result = transport.publish_dataset(PublishRequest(input_config="{}", data=pd.DataFrame({"x": [1]}))) + + assert transport.trace_id == "publish-trace" + assert not hasattr(result, "trace_id") + def test_status_parsed_from_json(self) -> None: transport = _make_flight_transport() _wire_do_put(transport, {**_VALID_JSON_RESP, "success": False}) diff --git a/tests/test_transport_trace_metadata.py b/tests/test_transport_trace_metadata.py index bb658d4..247b795 100644 --- a/tests/test_transport_trace_metadata.py +++ b/tests/test_transport_trace_metadata.py @@ -75,6 +75,19 @@ def test_get_ticket_returns_no_trace_id_when_schema_metadata_is_absent() -> None assert isinstance(result.schema_bytes, bytes) +def test_get_ticket_ignores_malformed_trace_id_metadata() -> None: + transport = _make_transport() + reader = MagicMock() + reader.schema = pa.schema([("id", pa.int64())], metadata={b"trace_id": b"\xff"}) + reader.read_chunk.side_effect = StopIteration + transport._client.do_get.return_value = reader + + result = transport.get_ticket(DatasetTicket(dataset_uuid="dataset-1")) + + assert isinstance(result.schema_bytes, bytes) + assert transport.trace_id is None + + def test_client_middleware_captures_response_trace_header_and_clears_previous_value() -> None: transport = _make_transport() transport._trace_id = "previous-trace"