|
| 1 | +from enum import Enum |
1 | 2 | from uuid import UUID |
2 | 3 |
|
3 | 4 | import pandas as pd |
|
25 | 26 | from tilebox.datasets.datasets.stac.v1.core_pb2 import Provider as ProviderPB2 |
26 | 27 | from tilebox.datasets.datasets.stac.v1.processing_pb import ProcessingSoftware |
27 | 28 | 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 | +) |
28 | 34 | from tilebox.datasets.datasets.stac.v1.storage_pb import Storage |
29 | 35 | from tilebox.datasets.datasets.stac.v1.storage_pb2 import Storage as StoragePB2 |
| 36 | +from tilebox.datasets.datasets.v1.well_known_types_pb2 import ProcessingLevel |
30 | 37 | from tilebox.datasets.protobuf_conversion.field_types import _AssetsDisplay |
31 | 38 | from tilebox.datasets.protobuf_conversion.protobuf_xarray import MessageToXarrayConverter |
32 | 39 | from tilebox.datasets.protobuf_conversion.to_protobuf import to_messages |
@@ -87,7 +94,7 @@ def test_convert_datapoint(datapoint: ExampleDatapoint) -> None: # noqa: PLR091 |
87 | 94 | ) |
88 | 95 |
|
89 | 96 | 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_") |
91 | 98 |
|
92 | 99 | assert list(dataset.some_repeated_string.to_numpy()) == list(datapoint.some_repeated_string) |
93 | 100 | 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: |
213 | 220 | assert "access_profiles" not in html |
214 | 221 |
|
215 | 222 |
|
| 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 | + |
216 | 306 | @given(lists(example_datapoints(generated_fields=True, missing_fields=True), min_size=5, max_size=30)) |
217 | 307 | def test_convert_datapoints(datapoints: list[ExampleDatapoint]) -> None: # noqa: C901, PLR0912 |
218 | 308 | converter = MessageToXarrayConverter() |
|
0 commit comments