Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
11 changes: 8 additions & 3 deletions dataconnect/client.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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(
Expand Down Expand Up @@ -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,
Expand Down
5 changes: 5 additions & 0 deletions dataconnect/exceptions.py
Original file line number Diff line number Diff line change
Expand Up @@ -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 = [
Expand All @@ -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)


Expand Down
23 changes: 15 additions & 8 deletions dataconnect/service/default.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand All @@ -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)
Expand All @@ -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).
Expand All @@ -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

Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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

Expand All @@ -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)
Expand Down
40 changes: 34 additions & 6 deletions dataconnect/service/error_handler.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
)
7 changes: 7 additions & 0 deletions dataconnect/transport/arrow_flight/error_handler.py
Original file line number Diff line number Diff line change
Expand Up @@ -205,13 +205,15 @@ 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(
error_code=error_code,
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_"):
Expand All @@ -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_"):
Expand All @@ -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_"):
Expand All @@ -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_"):
Expand All @@ -244,13 +249,15 @@ 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=error_code,
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)
Expand Down
Loading
Loading