Skip to content

Commit db59e9f

Browse files
Enums as strings in xarray (#51)
1 parent 33a1910 commit db59e9f

5 files changed

Lines changed: 125 additions & 42 deletions

File tree

‎CHANGELOG.md‎

Lines changed: 10 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -7,6 +7,14 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0
77

88
## [Unreleased]
99

10+
## [0.60.0] - 2026-08-25
11+
12+
### Changed
13+
14+
- `tilebox-datasets`: Dataset queries now return short enum names as strings (for example, `HH`) instead of numeric
15+
values with a name mapping in the xarray variable attributes. Repeated enum fields can now be queried and
16+
round-tripped through all supported ingestion inputs.
17+
1018
## [0.59.0] - 2026-08-18
1119

1220
### Added
@@ -471,7 +479,8 @@ the first client that does not cache data (since it's already on the local file
471479
- Released under the [MIT](https://opensource.org/license/mit) license.
472480
- Released packages: `tilebox-datasets`, `tilebox-workflows`, `tilebox-storage`, `tilebox-grpc`
473481

474-
[Unreleased]: https://github.com/tilebox/tilebox-python/compare/v0.59.0...HEAD
482+
[Unreleased]: https://github.com/tilebox/tilebox-python/compare/v0.60.0...HEAD
483+
[0.60.0]: https://github.com/tilebox/tilebox-python/compare/v0.59.0...v0.60.0
475484
[0.59.0]: https://github.com/tilebox/tilebox-python/compare/v0.58.0...v0.59.0
476485
[0.58.0]: https://github.com/tilebox/tilebox-python/compare/v0.57.0...v0.58.0
477486
[0.57.0]: https://github.com/tilebox/tilebox-python/compare/v0.56.0...v0.57.0

‎tilebox-datasets/tests/protobuf_conversion/test_protobuf_xarray.py‎

Lines changed: 91 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1,3 +1,4 @@
1+
from enum import Enum
12
from uuid import UUID
23

34
import pandas as pd
@@ -25,8 +26,14 @@
2526
from tilebox.datasets.datasets.stac.v1.core_pb2 import Provider as ProviderPB2
2627
from tilebox.datasets.datasets.stac.v1.processing_pb import ProcessingSoftware
2728
from tilebox.datasets.datasets.stac.v1.processing_pb2 import ProcessingSoftware as ProcessingSoftwarePB2
29+
from tilebox.datasets.datasets.stac.v1.sar_pb2 import (
30+
SAR_POLARIZATION_HH,
31+
SAR_POLARIZATION_VV,
32+
SARProperties,
33+
)
2834
from tilebox.datasets.datasets.stac.v1.storage_pb import Storage
2935
from tilebox.datasets.datasets.stac.v1.storage_pb2 import Storage as StoragePB2
36+
from tilebox.datasets.datasets.v1.well_known_types_pb2 import ProcessingLevel
3037
from tilebox.datasets.protobuf_conversion.field_types import _AssetsDisplay
3138
from tilebox.datasets.protobuf_conversion.protobuf_xarray import MessageToXarrayConverter
3239
from tilebox.datasets.protobuf_conversion.to_protobuf import to_messages
@@ -87,7 +94,7 @@ def test_convert_datapoint(datapoint: ExampleDatapoint) -> None: # noqa: PLR091
8794
)
8895

8996
assert isinstance(dataset.some_geometry.item(), Polygon | MultiPolygon)
90-
assert dataset.some_enum.item() == datapoint.some_enum
97+
assert dataset.some_enum.item() == ProcessingLevel.Name(datapoint.some_enum).removeprefix("PROCESSING_LEVEL_")
9198

9299
assert list(dataset.some_repeated_string.to_numpy()) == list(datapoint.some_repeated_string)
93100
assert_array_equal(dataset.some_repeated_int.to_numpy(), datapoint.some_repeated_int)
@@ -213,6 +220,89 @@ def test_convert_stac_messages_to_protobuf_py() -> None:
213220
assert "access_profiles" not in html
214221

215222

223+
def test_convert_scalar_and_repeated_enums_to_names_and_round_trip() -> None:
224+
class Polarization(Enum):
225+
HH = "HH"
226+
VV = "VV"
227+
228+
file_descriptor = descriptor_pb2.FileDescriptorProto(
229+
name="tests/protobuf_conversion/enum_datapoint.proto",
230+
package="tests.protobuf_conversion.enums",
231+
)
232+
enum_descriptor = file_descriptor.enum_type.add(name="SARPolarization")
233+
for name, number in (("SAR_POLARIZATION_UNSPECIFIED", 0), ("SAR_POLARIZATION_HH", 1), ("SAR_POLARIZATION_VV", 2)):
234+
enum_descriptor.value.add(name=name, number=number)
235+
message_descriptor = file_descriptor.message_type.add(name="EnumDatapoint")
236+
message_descriptor.field.add(
237+
name="primary_polarization",
238+
number=1,
239+
label=descriptor_pb2.FieldDescriptorProto.LABEL_OPTIONAL,
240+
type=descriptor_pb2.FieldDescriptorProto.TYPE_ENUM,
241+
type_name=".tests.protobuf_conversion.enums.SARPolarization",
242+
)
243+
message_descriptor.field.add(
244+
name="polarizations",
245+
number=2,
246+
label=descriptor_pb2.FieldDescriptorProto.LABEL_REPEATED,
247+
type=descriptor_pb2.FieldDescriptorProto.TYPE_ENUM,
248+
type_name=".tests.protobuf_conversion.enums.SARPolarization",
249+
)
250+
descriptor = Default().AddSerializedFile(file_descriptor.SerializeToString())
251+
message_type = GetMessageClass(descriptor.message_types_by_name["EnumDatapoint"])
252+
messages = [
253+
message_type(primary_polarization=1, polarizations=[0, 1]),
254+
message_type(primary_polarization=2, polarizations=[2]),
255+
message_type(),
256+
]
257+
258+
converter = MessageToXarrayConverter()
259+
converter.convert_all(messages)
260+
dataset = converter.finalize("time")
261+
262+
assert dataset.primary_polarization.dtype == object
263+
assert dataset.primary_polarization[:2].to_numpy().tolist() == ["HH", "VV"]
264+
assert pd.isna(dataset.primary_polarization[2].item())
265+
assert dataset.polarizations.dtype == object
266+
assert dataset.polarizations[0].to_numpy().tolist() == ["UNSPECIFIED", "HH"]
267+
assert dataset.polarizations[1, 0].item() == "VV"
268+
assert pd.isna(dataset.polarizations[1, 1].item())
269+
assert pd.isna(dataset.polarizations[2].to_numpy()).all()
270+
assert dataset.polarizations.dims == ("time", "n_polarizations")
271+
assert dataset.primary_polarization.attrs == {}
272+
assert dataset.polarizations.attrs == {}
273+
assert to_messages(dataset, message_type) == messages
274+
275+
expected = message_type(primary_polarization=1, polarizations=[1, 2, 0])
276+
record = {"primary_polarization": "HH", "polarizations": ["HH", Polarization.VV, 0]}
277+
assert to_messages([record], message_type) == [expected]
278+
assert to_messages({name: [value] for name, value in record.items()}, message_type) == [expected]
279+
assert to_messages(pd.DataFrame([record]), message_type) == [expected]
280+
281+
with pytest.raises(ValueError, match="Record 0: Field 'polarizations': Invalid enum name 'INVALID'"):
282+
to_messages([{"polarizations": ["INVALID"]}], message_type)
283+
284+
285+
def test_convert_sar_polarizations_with_short_names_in_both_directions() -> None:
286+
messages = [
287+
SARProperties(polarizations=[SAR_POLARIZATION_HH, SAR_POLARIZATION_VV]),
288+
SARProperties(polarizations=[SAR_POLARIZATION_VV]),
289+
SARProperties(),
290+
]
291+
converter = MessageToXarrayConverter()
292+
converter.convert_all(messages)
293+
294+
dataset = converter.finalize("item")
295+
296+
assert dataset.polarizations[0].to_numpy().tolist() == ["HH", "VV"]
297+
assert dataset.polarizations[1, 0].item() == "VV"
298+
assert pd.isna(dataset.polarizations[1, 1].item())
299+
assert pd.isna(dataset.polarizations[2].to_numpy()).all()
300+
301+
other_fields = [field.name for field in SARProperties.DESCRIPTOR.fields if field.name != "polarizations"]
302+
assert to_messages(dataset, SARProperties, ignore_fields=other_fields) == messages
303+
assert to_messages([{"polarizations": ["HH", "VV"]}], SARProperties) == [messages[0]]
304+
305+
216306
@given(lists(example_datapoints(generated_fields=True, missing_fields=True), min_size=5, max_size=30))
217307
def test_convert_datapoints(datapoints: list[ExampleDatapoint]) -> None: # noqa: C901, PLR0912
218308
converter = MessageToXarrayConverter()

‎tilebox-datasets/tests/test_client.py‎

Lines changed: 1 addition & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -145,8 +145,7 @@ def test_find_datapoint() -> None:
145145

146146
if not skip_data:
147147
assert datapoint.granule_name.item() == "S2A_MSIL1C_20220713T002201_N0400_R102_T08XNS_20220713T015332.SAFE"
148-
processing_level = datapoint.processing_level.item()
149-
assert datapoint.processing_level.attrs["names"][processing_level] == "L1C"
148+
assert datapoint.processing_level.item() == "L1C"
150149
assert datapoint.copernicus_id.item() == "65505f82-76dd-5e85-b947-a6c879e07446"
151150
assert isinstance(datapoint.geometry.item(), Polygon)
152151
else:

‎tilebox-datasets/tilebox/datasets/protobuf_conversion/field_types.py‎

Lines changed: 23 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -1,5 +1,7 @@
1+
import re
12
from collections.abc import Sequence
23
from datetime import timedelta
4+
from enum import Enum
35
from typing import Any
46
from uuid import UUID
57

@@ -120,21 +122,30 @@ def to_proto(self, value: Any) -> bool:
120122

121123
class EnumField(ProtobufFieldType):
122124
def __init__(self, name_lookup: dict[int, str]) -> None:
123-
super().__init__(np.uint8) # we support up to 256 different enum values for now
125+
super().__init__(object)
124126
self._values_to_name = name_lookup
125127
self._names_to_value = {name: value for value, name in name_lookup.items()}
126128

127-
def from_proto(self, value: ProtoFieldValue) -> int:
129+
def from_proto(self, value: ProtoFieldValue) -> str:
128130
if not isinstance(value, int):
129131
raise TypeError(f"Expected int message but got {type(value)}")
130-
return value # we don't parse the value when loading, to avoid having huge arrays of strings
131-
132-
def to_proto(self, value: str | int) -> int:
132+
try:
133+
return self._values_to_name[value]
134+
except KeyError as error:
135+
raise ValueError(f"Invalid enum value {value}") from error
136+
137+
def to_proto(self, value: str | int | Enum) -> int:
138+
if isinstance(value, Enum):
139+
value = value.name
133140
if isinstance(value, (str, np.str_)):
134-
return self._names_to_value[value]
135-
if int(value) not in self._values_to_name:
141+
try:
142+
return self._names_to_value[value]
143+
except KeyError as error:
144+
raise ValueError(f"Invalid enum name {value!r}") from error
145+
integer_value = int(value)
146+
if integer_value not in self._values_to_name:
136147
raise ValueError(f"Invalid enum value {value}") # during ingestion, we can raise an error here
137-
return value
148+
return integer_value
138149

139150

140151
class TimestampField(ProtobufFieldType):
@@ -360,8 +371,11 @@ def _camel_to_uppercase(name: str) -> str:
360371
Examples:
361372
>>> _camel_to_uppercase("ProcessingLevel")
362373
'PROCESSING_LEVEL'
374+
>>> _camel_to_uppercase("SARPolarization")
375+
'SAR_POLARIZATION'
363376
"""
364-
return "".join(["_" + c.lower() if c.isupper() else c for c in name]).lstrip("_").upper()
377+
name = re.sub(r"(.)([A-Z][a-z]+)", r"\1_\2", name)
378+
return re.sub(r"([a-z0-9])([A-Z])", r"\1_\2", name).upper()
365379

366380

367381
def is_missing(value: Any) -> bool:

‎tilebox-datasets/tilebox/datasets/protobuf_conversion/protobuf_xarray.py‎

Lines changed: 0 additions & 29 deletions
Original file line numberDiff line numberDiff line change
@@ -13,10 +13,8 @@
1313
from numpy.typing import NDArray
1414

1515
from tilebox.datasets.protobuf_conversion.field_types import (
16-
EnumField,
1716
ProtobufFieldType,
1817
ProtoFieldValue,
19-
enum_mapping_from_field_descriptor,
2018
infer_field_type,
2119
)
2220

@@ -313,26 +311,6 @@ def _resize(self) -> None:
313311
self._data = data
314312

315313

316-
class _EnumFieldConverter(_SimpleFieldConverter):
317-
def __init__(self, field_name: str, enum_names: dict[int, str]) -> None:
318-
"""
319-
A field converter for the enum type.
320-
321-
Args:
322-
field_name: The name of enum field in the protobuf message
323-
"""
324-
super().__init__(field_name, EnumField(enum_names))
325-
self._enum_names = enum_names
326-
327-
def finalize(
328-
self, dataset: xr.Dataset, count: int, dimension_names: tuple[str, ...], skip_if_empty: bool = False
329-
) -> str | None:
330-
field_name = super().finalize(dataset, count, dimension_names, skip_if_empty)
331-
if field_name is not None:
332-
dataset[field_name].attrs["names"] = self._enum_names
333-
return field_name
334-
335-
336314
def _create_field_converters(message: Message, buffer_size: int) -> dict[str, _FieldConverter]:
337315
"""
338316
Create a dictionary mapping from field names to field converters for the given protobuf message descriptor.
@@ -369,13 +347,6 @@ def _create_field_converter(field: FieldDescriptor) -> _FieldConverter:
369347
Returns:
370348
A field converter for the given protobuf field descriptor
371349
"""
372-
# special handling for enums:
373-
if field.type == FieldDescriptor.TYPE_ENUM:
374-
if field.is_repeated:
375-
raise NotImplementedError("Repeated enum fields are not supported")
376-
377-
return _EnumFieldConverter(field.name, enum_mapping_from_field_descriptor(field))
378-
379350
field_type = infer_field_type(field)
380351
if field.is_repeated:
381352
return _ArrayFieldConverter(field.name, field_type)

0 commit comments

Comments
 (0)