diff --git a/dataconnect/client.py b/dataconnect/client.py index 6acbf22..a26c151 100644 --- a/dataconnect/client.py +++ b/dataconnect/client.py @@ -34,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, @@ -67,7 +72,7 @@ def fetch_data( dataset_uuid: UUID, first_n_rows: int | None = None, ) -> pd.DataFrame: - """Fetch data frames for a given dataset UUID.""" + """Fetch data for a given dataset UUID.""" return self._service.fetch_data(dataset_uuid, first_n_rows) def get_datasets( @@ -121,8 +126,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. """ return self._service.dry_publish( 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/service/default.py b/dataconnect/service/default.py index 51a2845..c7bb9cf 100644 --- a/dataconnect/service/default.py +++ b/dataconnect/service/default.py @@ -74,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: @@ -92,7 +96,8 @@ 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 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) @@ -106,7 +111,7 @@ def get_dataset_versions(self, dataset_uuid: UUID) -> list[DatasetVersion]: dataset_uuid: UUID of the dataset whose versions are requested. Returns: - A list of :class:`DatasetVersion` objects for the given dataset. + The dataset versions, newest first. Raises: ValidationError: If *dataset_uuid* is not a valid UUID (upstream). @@ -115,14 +120,15 @@ 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), key=lambda dv: dv.dataset_version, reverse=True, ) + return versions except Exception as ex: raise translate_error(ex) from ex @@ -169,7 +175,8 @@ def get_datasets( ) try: - resources = self._transport.list_resources(request) + result = self._transport.list_resources(request) + resources = result items = [] for resource in resources: dataset = resource_to_dataset(resource) @@ -347,7 +354,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,7 +363,7 @@ 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 ] return DatetimeFormatsResult(formats=formats) 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/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 7dfe1ff..f76b2d2 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 @@ -87,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", @@ -116,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}" @@ -172,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 @@ -182,6 +249,21 @@ 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) -> list[ResourceInfo]: @@ -200,7 +282,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()) + raw_flights = list(self._client.list_flights(criteria, self._options())) + + if raw_flights and raw_flights[0].app_metadata: + try: + 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): + pass + return [_to_resource_info(f) for f in raw_flights] except Exception as ex: @@ -224,6 +315,16 @@ def get_ticket(self, ticket: DatasetTicket) -> DataTable: except flight.FlightError as ex: raise parse_dataconnect_error(ex) from ex + if table.schema.metadata: + raw_trace_id = table.schema.metadata.get(b"trace_id") + if raw_trace_id is not None: + 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)) except Exception as ex: @@ -283,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() @@ -340,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() @@ -384,7 +487,13 @@ def get_datetime_formats(self, request: DatetimeFormatsRequest) -> list[str]: if not results: return [] - return json.loads(results[0].body.to_pybytes().decode("utf-8")) + response = json.loads(results[0].body.to_pybytes().decode("utf-8")) + 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 91ce61f..1ce7078 100644 --- a/dataconnect/transport/base.py +++ b/dataconnect/transport/base.py @@ -30,6 +30,9 @@ def list_resources(self, request: ResourceQuery) -> list[ResourceInfo]: The transport does not interpret the action name or body — that is the service layer's responsibility. + + Returns: + The matched resources. """ @abstractmethod @@ -63,6 +66,9 @@ def get_datetime_formats(self, request: DatetimeFormatsRequest) -> list[str]: 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: + 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/tests/test_dry_publish.py b/tests/test_dry_publish.py index 45d4dd6..86ca44d 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,6 +30,7 @@ from dataconnect.transport.models import ( DatasetTicket, DataTable, + DatetimeFormatsRequest, DryPublishResponse, PublishRequest, PublishResponse, @@ -225,6 +226,10 @@ 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_is_not_added_to_result_contract(self) -> None: + result = dry_publish_response_to_domain(_make_dry_publish_response()) + assert not hasattr(result, "trace_id") + # --------------------------------------------------------------------------- # DefaultDataConnectService.dry_publish @@ -258,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) -> list[str]: # type: ignore[override] + def get_datetime_formats(self, request: DatetimeFormatsRequest) -> list[str]: raise NotImplementedError def close(self) -> None: @@ -514,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_fetch_data.py b/tests/test_fetch_data.py index a431386..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, ResourceInfo, ResourceQuery +from dataconnect.transport.models import DatasetTicket, DataTable, ResourceQuery # --------------------------------------------------------------------------- # Fake transport @@ -41,7 +41,7 @@ def __init__( self._get_ticket_error = get_ticket_error self.last_ticket: DatasetTicket | None = None - def list_resources(self, request: ResourceQuery) -> list[ResourceInfo]: + def list_resources(self, request: ResourceQuery) -> list: return [] def get_ticket(self, ticket: DatasetTicket) -> DataTable: diff --git a/tests/test_get_datasets_paginated.py b/tests/test_get_datasets_paginated.py index d3cae99..9af6d24 100644 --- a/tests/test_get_datasets_paginated.py +++ b/tests/test_get_datasets_paginated.py @@ -16,7 +16,13 @@ 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, + ResourceQuery, +) class _FakeTransport: @@ -24,9 +30,11 @@ 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]: @@ -156,6 +164,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) + + service.get_datasets(study_environment_uuid=_STUDY_ENV_UUID) + + assert service.trace_id == "trace-datasets-1" + + def test_trace_id_defaults_to_none(self) -> None: + transport = _FakeTransport(resources=[]) + service = DefaultDataConnectService(transport) + + service.get_datasets(study_environment_uuid=_STUDY_ENV_UUID) + + assert service.trace_id is None + def test_multiple_items_returned(self) -> None: resources = [ _dataset_resource( @@ -231,7 +255,8 @@ 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()) - 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 resources def get_ticket(self, ticket: DatasetTicket) -> DataTable: if self.closed: @@ -477,3 +502,15 @@ def test_dataset_versions_response_is_unchanged_before_and_after_listing() -> No 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) + + service.get_dataset_versions(UUID(_DATASET_UUID)) + + assert service.trace_id == "trace-versions-1" diff --git a/tests/test_get_datetime_formats.py b/tests/test_get_datetime_formats.py index 5a754c2..7bfbd49 100644 --- a/tests/test_get_datetime_formats.py +++ b/tests/test_get_datetime_formats.py @@ -28,7 +28,6 @@ DryPublishResponse, PublishRequest, PublishResponse, - ResourceInfo, ResourceQuery, ) @@ -52,12 +51,14 @@ 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]: + def list_resources(self, request: ResourceQuery) -> list: return [] def get_ticket(self, ticket: DatasetTicket) -> DataTable: @@ -82,8 +83,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 +251,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") + 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() + service.get_datetime_formats(project_token="tok") + assert service.trace_id is None + # --- error translation --- def test_transport_validation_error_is_translated_to_service_error(self) -> None: @@ -274,9 +286,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]) -> MagicMock: body = MagicMock() - body.to_pybytes.return_value = json.dumps(payload).encode("utf-8") + body.to_pybytes.return_value = json.dumps(formats).encode("utf-8") result = MagicMock() result.body = body return result @@ -310,17 +322,17 @@ 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 == _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 result == [] + 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 35e531d..4a9b342 100644 --- a/tests/test_publish.py +++ b/tests/test_publish.py @@ -146,6 +146,10 @@ def test_returns_publish_result_instance(self) -> None: result = publish_response_to_domain(_make_publish_response()) assert isinstance(result, PublishResult) + def test_trace_id_is_not_added_to_result_contract(self) -> None: + result = publish_response_to_domain(_make_publish_response()) + assert not hasattr(result, "trace_id") + # --------------------------------------------------------------------------- # DefaultDataConnectService.publish @@ -179,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) -> list[str]: # type: ignore[override] + def get_datetime_formats(self, request: DatetimeFormatsRequest) -> list[str]: raise NotImplementedError def close(self) -> None: @@ -435,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 new file mode 100644 index 0000000..247b795 --- /dev/null +++ b/tests/test_transport_trace_metadata.py @@ -0,0 +1,129 @@ +"""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.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 + + +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 transport.trace_id == "trace-list-1" + assert len(result) == 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)] + + transport._begin_call() + result = transport.list_resources(ResourceQuery(action="studies.list")) + + assert transport.trace_id is None + assert isinstance(result, list) + + +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 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: + 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 transport.trace_id is 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" + + 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" + )