From d0fde3e9e5a1343caa8d6697ef6008faa1825cf9 Mon Sep 17 00:00:00 2001 From: awarde96 Date: Fri, 19 Jun 2026 09:24:14 +0000 Subject: [PATCH 1/5] Add identifiers to mars:metadata if they are passed from polytope --- covjsonkit/encoder/TimeSeries.py | 24 +++++++++++++++++------- covjsonkit/encoder/encoder.py | 12 ++++++++++++ 2 files changed, 29 insertions(+), 7 deletions(-) diff --git a/covjsonkit/encoder/TimeSeries.py b/covjsonkit/encoder/TimeSeries.py index 310f630..699e6ed 100644 --- a/covjsonkit/encoder/TimeSeries.py +++ b/covjsonkit/encoder/TimeSeries.py @@ -143,6 +143,7 @@ def from_polytope(self, result, date_key: str = "date") -> dict: fields["step"] = 0 fields["dates"] = [] fields["levels"] = [0] + fields["identifiers"] = [] start = time.time() logging.debug("Tree walking starts at: %s", start) # noqa: E501 @@ -243,13 +244,22 @@ def from_polytope(self, result, date_key: str = "date") -> dict: f"Key {key} not found in range_dict. " f"Please ensure all axes are compressed in config" ) - mm = mars_metadata.copy() - mm["number"] = num - mm["Forecast date"] = date - mm["levelist"] = level - coordinates[date][i]["levelist"] = [level] - del mm["step"] - self.add_coverage(mm, coordinates[date][i], val_dict) + # Determine identifiers for this point + point_identifiers = [None] + if fields["identifiers"] and i < len(fields["identifiers"]): + point_identifiers = fields["identifiers"][i] + + # Emit one coverage per identifier (duplicates coverage for merged points) + for identifier in point_identifiers: + mm = mars_metadata.copy() + mm["number"] = num + mm["Forecast date"] = date + mm["levelist"] = level + if identifier is not None: + mm["identifier"] = identifier + coordinates[date][i]["levelist"] = [level] + del mm["step"] + self.add_coverage(mm, coordinates[date][i], val_dict) end = time.time() delta = end - start diff --git a/covjsonkit/encoder/encoder.py b/covjsonkit/encoder/encoder.py index 0332de2..c554f7f 100644 --- a/covjsonkit/encoder/encoder.py +++ b/covjsonkit/encoder/encoder.py @@ -383,6 +383,17 @@ def append_composite_coords(dates, tree_values, lat, coords): for value in tree_values: coords[dates]["composite"].append([lat, value]) + def collect_tags(tree, fields, count): + """Collect tags from leaf tree nodes into fields['identifiers'].""" + if "identifiers" in fields: + tags = getattr(tree, "tags", None) + if tags: + tag_list = sorted(tags, key=str) + else: + tag_list = [None] + for _ in range(count): + fields["identifiers"].append(tag_list) + if len(tree.children) != 0: for child in tree.children: handle_non_leaf_node(child) @@ -426,6 +437,7 @@ def append_composite_coords(dates, tree_values, lat, coords): step_len = para_len / len(fields["step"]) append_composite_coords(fields["dates"][-1], tree.values, fields["lat"], coords) + collect_tags(tree, fields, len(tree.values)) for l, level in enumerate(fields["levels"]): # noqa: E741 for i, num in enumerate(fields["number"]): From 92a0eafd5158b1bb65583b27e98de83c7c1937db Mon Sep 17 00:00:00 2001 From: awarde96 Date: Tue, 23 Jun 2026 10:19:01 +0000 Subject: [PATCH 2/5] Take tags from latitude instead of longitude for labels --- covjsonkit/encoder/encoder.py | 5 +++-- 1 file changed, 3 insertions(+), 2 deletions(-) diff --git a/covjsonkit/encoder/encoder.py b/covjsonkit/encoder/encoder.py index c554f7f..2bcafcd 100644 --- a/covjsonkit/encoder/encoder.py +++ b/covjsonkit/encoder/encoder.py @@ -384,9 +384,9 @@ def append_composite_coords(dates, tree_values, lat, coords): coords[dates]["composite"].append([lat, value]) def collect_tags(tree, fields, count): - """Collect tags from leaf tree nodes into fields['identifiers'].""" + """Collect tags from the latitude node into fields['identifiers'].""" if "identifiers" in fields: - tags = getattr(tree, "tags", None) + tags = fields.get("_lat_tags", None) or getattr(tree, "tags", None) if tags: tag_list = sorted(tags, key=str) else: @@ -401,6 +401,7 @@ def collect_tags(tree, fields, count): if result is not None: if child.axis.name == "latitude": fields["lat"] = result + fields["_lat_tags"] = getattr(child, "tags", None) elif child.axis.name == "levelist": fields["levels"] = result if "l" in fields: From dfceb57bef27360766268ae0feb65d5fb3e97ed2 Mon Sep 17 00:00:00 2001 From: awarde96 Date: Wed, 24 Jun 2026 10:05:42 +0000 Subject: [PATCH 3/5] Change output in metadata from identifier to label --- covjsonkit/encoder/TimeSeries.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/covjsonkit/encoder/TimeSeries.py b/covjsonkit/encoder/TimeSeries.py index 699e6ed..4ed7395 100644 --- a/covjsonkit/encoder/TimeSeries.py +++ b/covjsonkit/encoder/TimeSeries.py @@ -256,7 +256,7 @@ def from_polytope(self, result, date_key: str = "date") -> dict: mm["Forecast date"] = date mm["levelist"] = level if identifier is not None: - mm["identifier"] = identifier + mm["label"] = identifier coordinates[date][i]["levelist"] = [level] del mm["step"] self.add_coverage(mm, coordinates[date][i], val_dict) From 5ad928af9adfbda598d3622b2baf1d0a5b3e7857 Mon Sep 17 00:00:00 2001 From: awarde96 Date: Thu, 1 Oct 2026 07:40:07 +0000 Subject: [PATCH 4/5] Use new label format passed from polytope for point like data --- covjsonkit/encoder/Position.py | 42 ++- covjsonkit/encoder/TimeSeries.py | 134 ++++++---- covjsonkit/encoder/VerticalProfile.py | 33 ++- covjsonkit/encoder/encoder.py | 123 +++++++-- tests/test_encoder_labels.py | 361 ++++++++++++++++++++++++++ 5 files changed, 605 insertions(+), 88 deletions(-) create mode 100644 tests/test_encoder_labels.py diff --git a/covjsonkit/encoder/Position.py b/covjsonkit/encoder/Position.py index 819f03e..d56fc7c 100644 --- a/covjsonkit/encoder/Position.py +++ b/covjsonkit/encoder/Position.py @@ -4,7 +4,14 @@ import pandas as pd -from .encoder import Encoder, normalize_step_value +from .encoder import ( + Encoder, + add_label, + expand_points_by_tags, + expand_tags, + normalize_step_value, + tag_sort_key, +) class Position(Encoder): @@ -222,7 +229,9 @@ def from_polytope(self, result, date_key: str = "date") -> dict: logging.debug("The fields retrieved were: %s", fields) # noqa: E501 logging.debug("The range_dict created was: %s", range_dict) # noqa: E501 - for i, point in enumerate(range(points)): + entries = expand_points_by_tags(coords[fields["dates"][0]].get("tags"), points) + + for i, label in entries: for date in fields["dates"]: for level in fields["levels"]: for num in fields["number"]: @@ -246,6 +255,7 @@ def from_polytope(self, result, date_key: str = "date") -> dict: mm["number"] = num mm["Forecast date"] = date del mm["step"] + add_label(mm, label) self.add_coverage(mm, coordinates[date][i], val_dict) end = time.time() @@ -293,7 +303,11 @@ def from_polytope_reforecast(self, result) -> dict: coverage_order = [] param_order = [] - for rec in self._reforecast_records(result): + for rec, tag_index, label in ( + (rec, tag_index, label) + for rec in self._reforecast_records(result) + for tag_index, label in expand_tags(rec["__tags__"]) + ): value = float(rec["__value__"]) lat = float(rec["latitude"]) lon = float(rec["longitude"]) @@ -313,15 +327,16 @@ def from_polytope_reforecast(self, result) -> dict: valid = ref + self._reforecast_step_timedelta(step) valid_iso = valid.isoformat() + "Z" - key = (lat, lon, level, number, ref.isoformat()) + key = (lat, lon, level, number, ref.isoformat(), tag_index) if key not in coverages: meta = {} for name in rec: - if name == "__value__" or name in exclude_meta: + if name in ("__value__", "__tags__") or name in exclude_meta: continue meta[name] = self._reforecast_stringify(rec[name]) meta["number"] = number meta["Forecast date"] = ref.isoformat() + "Z" + add_label(meta, label) coverages[key] = { "lat": lat, "lon": lon, @@ -344,6 +359,9 @@ def from_polytope_reforecast(self, result) -> dict: for para in param_order: self.add_parameter(para) + # Tagged points: emit in request order (stable, so tree order is kept per point). + coverage_order.sort(key=lambda k: tag_sort_key(k[-1])) + for key in coverage_order: cov = coverages[key] times = sorted(cov["times"].keys()) @@ -406,8 +424,10 @@ def from_polytope_month(self, result): points = len(coords[fields["dates"][0]]["composite"]) # Build one coordinate entry per point per level; the t axis is all months. + entries = expand_points_by_tags(coords[fields["dates"][0]].get("tags"), points) + coordinates = [] - for i in range(points): + for i, label in entries: for level in fields["levels"]: coord_entry = { "latitude": [coords[fields["dates"][0]]["composite"][i][0]], @@ -415,7 +435,7 @@ def from_polytope_month(self, result): "levelist": [level], "t": [f"{date}-01T00:00:00Z" for date in fields["dates"]], } - coordinates.append((i, level, coord_entry)) + coordinates.append((i, label, level, coord_entry)) end = time.time() logging.debug("Coords creation: %s", end) # noqa: E501 @@ -428,7 +448,7 @@ def from_polytope_month(self, result): logging.debug("The fields retrieved were: %s", fields) # noqa: E501 logging.debug("The range_dict created was: %s", range_dict) # noqa: E501 - for i, level, coord_entry in coordinates: + for i, label, level, coord_entry in coordinates: for num in fields["number"]: val_dict = {} for para in fields["param"]: @@ -448,6 +468,7 @@ def from_polytope_month(self, result): mm = mars_metadata.copy() mm["number"] = num mm["levelist"] = level + add_label(mm, label) self.add_coverage(mm, coord_entry, val_dict) end = time.time() @@ -537,7 +558,9 @@ def from_polytope_step(self, result): start = time.time() logging.debug("Coverage creation: %s", start) # noqa: E501 - for i, point in enumerate(range(points)): + entries = expand_points_by_tags(coords[fields["dates"][0]].get("tags"), points) + + for i, label in entries: for j, level in enumerate(fields["levels"]): for num in fields["number"]: val_dict = {} @@ -552,6 +575,7 @@ def from_polytope_step(self, result): mm = mars_metadata.copy() mm["number"] = num mm["Forecast date"] = date + add_label(mm, label) self.add_coverage(mm, coordinates[fields["dates"][0]][(i * len(fields["levels"]) + j)], val_dict) end = time.time() diff --git a/covjsonkit/encoder/TimeSeries.py b/covjsonkit/encoder/TimeSeries.py index d6a6c89..7f25842 100644 --- a/covjsonkit/encoder/TimeSeries.py +++ b/covjsonkit/encoder/TimeSeries.py @@ -4,7 +4,14 @@ import pandas as pd -from .encoder import Encoder +from .encoder import ( + Encoder, + add_label, + expand_points_by_tags, + expand_tags, + node_tags, + tag_sort_key, +) class TimeSeries(Encoder): @@ -110,7 +117,9 @@ def _collapse_reanalysis(self, fields, coords, mars_metadata, range_dict): t_values = [stamp.isoformat() + "Z" for stamp, _, _ in stamp_order] - for i in range(points): + entries = expand_points_by_tags(coords[first_date].get("tags"), points) + + for i, label in entries: lat = coords[first_date]["composite"][i][0] lon = coords[first_date]["composite"][i][1] for level in fields["levels"]: @@ -132,6 +141,7 @@ def _collapse_reanalysis(self, fields, coords, mars_metadata, range_dict): mm["levelist"] = level mm.pop("step", None) mm.pop("Forecast date", None) + add_label(mm, label) coord_entry = { "latitude": [lat], "longitude": [lon], @@ -323,7 +333,9 @@ def from_polytope(self, result, date_key: str = "date", reforecast: bool = False logging.debug("The fields retrieved were: %s", fields) # noqa: E501 logging.debug("The range_dict created was: %s", range_dict) # noqa: E501 - for i, point in enumerate(range(points)): + entries = expand_points_by_tags(coords[fields["dates"][0]].get("tags"), points) + + for i, label in entries: for date in fields["dates"]: for level in fields["levels"]: for num in fields["number"]: @@ -349,6 +361,7 @@ def from_polytope(self, result, date_key: str = "date", reforecast: bool = False mm["levelist"] = level coordinates[date][i]["levelist"] = [level] del mm["step"] + add_label(mm, label) self.add_coverage(mm, coordinates[date][i], val_dict) end = time.time() @@ -426,7 +439,7 @@ def stringify(value): return str(value) return value - def emit(full_path, flat_result): + def emit(full_path, flat_result, tags): nonlocal forecast_result axis_names = [name for name, _ in full_path] axis_values = [values for _, values in full_path] @@ -467,55 +480,69 @@ def emit(full_path, flat_result): is_ce = d.get("class") == "ce" collapse = is_ce and stream == "efcl" forecast = is_ce and stream == "efas" - if collapse: - key = (lat, lon, level, number) - elif forecast: + if forecast: # Reference datetime of the forecast run = date + time. reference = self._hdate_step_timestamp(hdate, 0, time_off) - key = (lat, lon, level, number, reference) - else: - key = (lat, lon, level, number, stringify(hdate)) - if key not in coverages: - meta = {} - for name in axis_names: - if name in exclude_meta: - continue - meta[name] = stringify(d[name]) - meta["number"] = number - meta["levelist"] = level - if forecast: - # date+time is folded into the run reference; drop the - # raw date axis so it doesn't duplicate "Forecast date". - meta.pop("date", None) - meta["Forecast date"] = reference.isoformat() + "Z" - elif not collapse: - meta["Forecast date"] = pd.Timestamp(hdate).isoformat() + "Z" - coverages[key] = { - "lat": lat, - "lon": lon, - "level": level, - "number": number, - "meta": meta, - "params": {}, - } - if forecast: - forecast_result = True - # Ordering key so coverages come out date -> time -> point - # (reference = date + time). Insertion/tree order is - # otherwise time-major for multi-date forecast requests. - coverages[key]["sort_key"] = (reference, lat, lon, level, number) - coverage_order.append(key) - - if para not in param_order: - param_order.append(para) - params = coverages[key]["params"] - params.setdefault(para, []).append((stamp, float(value))) + + # One coverage per requested point: a grid point carrying several + # (index, label) tags is emitted once per tag. + for tag_index, label in expand_tags(tags): + if collapse: + key = (lat, lon, level, number, tag_index) + elif forecast: + key = (lat, lon, level, number, reference, tag_index) + else: + key = (lat, lon, level, number, stringify(hdate), tag_index) + if key not in coverages: + meta = {} + for name in axis_names: + if name in exclude_meta: + continue + meta[name] = stringify(d[name]) + meta["number"] = number + meta["levelist"] = level + if forecast: + # date+time is folded into the run reference; drop the + # raw date axis so it doesn't duplicate "Forecast date". + meta.pop("date", None) + meta["Forecast date"] = reference.isoformat() + "Z" + elif not collapse: + meta["Forecast date"] = pd.Timestamp(hdate).isoformat() + "Z" + add_label(meta, label) + coverages[key] = { + "lat": lat, + "lon": lon, + "level": level, + "number": number, + "meta": meta, + "params": {}, + "tag_index": tag_index, + } + if forecast: + forecast_result = True + # Ordering key so coverages come out date -> time -> point + # (reference = date + time). Insertion/tree order is + # otherwise time-major for multi-date forecast requests. + coverages[key]["sort_key"] = ( + reference, + tag_sort_key(tag_index), + lat, + lon, + level, + number, + ) + coverage_order.append(key) + + if para not in param_order: + param_order.append(para) + params = coverages[key]["params"] + params.setdefault(para, []).append((stamp, float(value))) def recurse(node, path): children = node.children if len(children) == 0: # Leaf longitude node: ``path`` already carries latitude/longitude. - emit(path, node.result) + emit(path, node.result, node_tags(node)) return for child in children: if is_merged_node(child): @@ -525,7 +552,7 @@ def recurse(node, path): ("latitude", (lat,)), ("longitude", (lon,)), ] - emit(merged_path, child.result) + emit(merged_path, child.result, node_tags(child)) continue recurse(child, path + [(child.axis.name, tuple(child.values))]) @@ -540,6 +567,9 @@ def recurse(node, path): # Forecast (efas) coverages: emit in date -> time -> point order. if forecast_result: coverage_order.sort(key=lambda k: coverages[k]["sort_key"]) + elif any(coverages[k]["tag_index"] is not None for k in coverage_order): + # Tagged points: emit in request order (stable, so tree order is kept per point). + coverage_order.sort(key=lambda k: tag_sort_key(coverages[k]["tag_index"])) for key in coverage_order: cov = coverages[key] @@ -638,7 +668,9 @@ def from_polytope_month(self, result): logging.debug("The fields retrieved were: %s", fields) logging.debug("The range_dict created was: %s", range_dict) - for i in range(points): + entries = expand_points_by_tags(coords[first_date].get("tags"), points) + + for i, label in entries: for j, level in enumerate(fields["levels"]): for num in fields["number"]: val_dict = {} @@ -659,6 +691,7 @@ def from_polytope_month(self, result): mm = mars_metadata.copy() mm["number"] = num mm["levelist"] = level + add_label(mm, label) # Use all date keys as the time series for this coverage. coord_entry = coordinates[first_date][i].copy() coord_entry["levelist"] = [level] @@ -752,7 +785,9 @@ def from_polytope_step(self, result): start = time.time() logging.debug("Coverage creation: %s", start) # noqa: E501 - for i, point in enumerate(range(points)): + entries = expand_points_by_tags(coords[fields["dates"][0]].get("tags"), points) + + for i, label in entries: for j, level in enumerate(fields["levels"]): for num in fields["number"]: val_dict = {} @@ -767,6 +802,7 @@ def from_polytope_step(self, result): mm = mars_metadata.copy() mm["number"] = num mm["Forecast date"] = date + add_label(mm, label) self.add_coverage(mm, coordinates[fields["dates"][0]][(i * len(fields["levels"]) + j)], val_dict) end = time.time() diff --git a/covjsonkit/encoder/VerticalProfile.py b/covjsonkit/encoder/VerticalProfile.py index 70cc0e5..e931797 100644 --- a/covjsonkit/encoder/VerticalProfile.py +++ b/covjsonkit/encoder/VerticalProfile.py @@ -4,7 +4,14 @@ import pandas as pd -from .encoder import Encoder, normalize_step_value +from .encoder import ( + Encoder, + add_label, + expand_points_by_tags, + expand_tags, + normalize_step_value, + tag_sort_key, +) class VerticalProfile(Encoder): @@ -204,7 +211,9 @@ def from_polytope(self, result, date_key: str = "date") -> dict: logging.debug("The fields retrieved were: %s", fields) # noqa: E501 logging.debug("The range_dict created was: %s", range_dict) # noqa: E501 - for i, point in enumerate(range(points)): + entries = expand_points_by_tags(coords[fields["dates"][0]].get("tags"), points) + + for i, label in entries: for date in fields["dates"]: for num in fields["number"]: val_dict = {} @@ -230,6 +239,7 @@ def from_polytope(self, result, date_key: str = "date") -> dict: mm["Forecast date"] = date mm["step"] = normalize_step_value(step) # del mm["step"] + add_label(mm, label) self.add_coverage(mm, coordinates[date][i][s], val_dict[step]) end = time.time() @@ -277,7 +287,11 @@ def from_polytope_reforecast(self, result) -> dict: coverage_order = [] param_order = [] - for rec in self._reforecast_records(result): + for rec, tag_index, label in ( + (rec, tag_index, label) + for rec in self._reforecast_records(result) + for tag_index, label in expand_tags(rec["__tags__"]) + ): value = float(rec["__value__"]) lat = float(rec["latitude"]) lon = float(rec["longitude"]) @@ -297,16 +311,17 @@ def from_polytope_reforecast(self, result) -> dict: valid = ref + self._reforecast_step_timedelta(step) valid_iso = valid.isoformat() + "Z" - key = (lat, lon, number, ref.isoformat(), str(step)) + key = (lat, lon, number, ref.isoformat(), str(step), tag_index) if key not in coverages: meta = {} for name in rec: - if name == "__value__" or name in exclude_meta: + if name in ("__value__", "__tags__") or name in exclude_meta: continue meta[name] = self._reforecast_stringify(rec[name]) meta["number"] = number meta["step"] = step meta["Forecast date"] = ref.isoformat() + "Z" + add_label(meta, label) coverages[key] = { "lat": lat, "lon": lon, @@ -331,6 +346,9 @@ def from_polytope_reforecast(self, result) -> dict: for para in param_order: self.add_parameter(para) + # Tagged points: emit in request order (stable, so tree order is kept per point). + coverage_order.sort(key=lambda k: tag_sort_key(k[-1])) + for key in coverage_order: cov = coverages[key] levels = sorted(cov["levels"].keys()) @@ -417,7 +435,9 @@ def from_polytope_month(self, result): logging.debug("The fields retrieved were: %s", fields) # noqa: E501 logging.debug("The range_dict created was: %s", range_dict) # noqa: E501 - for i in range(points): + entries = expand_points_by_tags(coords[fields["dates"][0]].get("tags"), points) + + for i, label in entries: for date in fields["dates"]: for num in fields["number"]: val_dict = {} @@ -438,6 +458,7 @@ def from_polytope_month(self, result): mm = mars_metadata.copy() mm["number"] = num mm["Forecast date"] = date + add_label(mm, label) self.add_coverage(mm, coordinates[date][i], val_dict) end = time.time() diff --git a/covjsonkit/encoder/encoder.py b/covjsonkit/encoder/encoder.py index 6f74b3f..ba4c662 100644 --- a/covjsonkit/encoder/encoder.py +++ b/covjsonkit/encoder/encoder.py @@ -1,5 +1,6 @@ from __future__ import annotations +import logging from abc import ABC, abstractmethod from datetime import datetime, timedelta from typing import Any @@ -32,6 +33,73 @@ def is_merged_node(node) -> bool: return hasattr(node, "axes") and getattr(node, "axes", None) is not None +def parse_tag(tag) -> tuple: + """Split a polytope shape tag into ``(index, label)``. + + polytope-mars tags each requested point (or polygon) as ``(index, label)``, where + ``index`` is its position in the request and ``label`` the user's label or ``None``. + Any other tag is treated as a bare label with no index. + """ + if isinstance(tag, tuple) and len(tag) == 2 and isinstance(tag[0], int) and not isinstance(tag[0], bool): + return tag + return (None, tag) + + +def node_tags(node) -> list: + """The ``(index, label)`` tags of a leaf node, sorted by index; empty if untagged.""" + tags = getattr(node, "tags", None) or () + parsed = [parse_tag(tag) for tag in tags] + return sorted(parsed, key=lambda t: (t[0] is None, t[0] if t[0] is not None else 0, str(t[1]))) + + +def expand_points_by_tags(tags_per_point, n_points) -> list: + """Return the ``(point_index, label)`` pairs to emit, in request order. + + ``tags_per_point`` holds the leaf tags of each spatial point, parallel to the + composite coordinates. Each tag gives one entry, so a grid point carrying two tags + (two requested points snapped to it) is emitted twice. Tagged entries come first, + ordered by request index; untagged points follow once each, without a label, in + tree order. A tree without tags therefore gives the same output as before. + """ + if tags_per_point is None or len(tags_per_point) != n_points: + if tags_per_point: + logging.warning("Found %s tag entries for %s points; ignoring tags", len(tags_per_point), n_points) + return [(i, None) for i in range(n_points)] + + tagged = [] + untagged = [] + for i, tags in enumerate(tags_per_point): + if not tags: + untagged.append((i, None)) + continue + for index, label in tags: + if index is None: + untagged.append((i, label)) + else: + tagged.append((index, i, label)) + if tagged and untagged: + logging.warning("%s of %s points carry no tag and are returned without a label", len(untagged), n_points) + tagged.sort(key=lambda t: (t[0], t[1])) + return [(i, label) for _, i, label in tagged] + untagged + + +def expand_tags(tags) -> list: + """``(index, label)`` pairs for one leaf value: one per tag, or a single untagged entry.""" + return list(tags) if tags else [(None, None)] + + +def tag_sort_key(index): + """Sort key placing tagged entries in request order before untagged ones.""" + return (index is None, index if index is not None else 0) + + +def add_label(metadata, label): + """Set ``label`` in coverage metadata when one was requested.""" + if label is not None: + metadata["label"] = label + return metadata + + def timedelta_to_step_string(td: timedelta) -> str: """ Convert a timedelta object to a step string in the format 'XhYm'. @@ -397,12 +465,13 @@ def calculate_index_bounds(level_len, num_len, para_len, step_len, l, i, j, k): end_index = start_index + int(step_len) return start_index, end_index - def append_composite_coords(dates, tree_values, lat, coords): + def append_composite_coords(dates, tree_values, lat, coords, tags): # for date in dates: for value in tree_values: coords[dates]["composite"].append([lat, value]) + coords[dates].setdefault("tags", []).append(tags) - def emit_leaf(lat, lon_values, result): + def emit_leaf(lat, lon_values, result, tags): """Emit one spatial leaf: append [lat, lon] composite coords and slice results. Shared by the legacy longitude-leaf path (``lat`` from the parent latitude @@ -427,7 +496,7 @@ def emit_leaf(lat, lon_values, result): para_len = num_len / len(fields["param"]) step_len = para_len / len(fields["step"]) - append_composite_coords(fields["dates"][-1], lon_values, lat, coords) + append_composite_coords(fields["dates"][-1], lon_values, lat, coords, tags) for l, level in enumerate(fields["levels"]): # noqa: E741 for i, num in enumerate(fields["number"]): @@ -445,7 +514,7 @@ def emit_leaf(lat, lon_values, result): for child in tree.children: # Compacted unstructured leaf: values=(lat, lon), own result. Emit directly. if is_merged_node(child): - emit_leaf(child.values[0], [child.values[1]], child.result) + emit_leaf(child.values[0], [child.values[1]], child.result, node_tags(child)) continue handle_non_leaf_node(child) result = handle_specific_axes(child) @@ -469,7 +538,7 @@ def emit_leaf(lat, lon_values, result): self.walk_tree(child, fields, coords, mars_metadata, range_dict, date_key=date_key) else: - emit_leaf(fields["lat"], tree.values, tree.result) + emit_leaf(fields["lat"], tree.values, tree.result, node_tags(tree)) def walk_tree_reforecast(self, tree, fields, coords, mars_metadata, range_dict): """Walk the result tree for reforecast/reanalysis with an independent ``time`` axis. @@ -536,11 +605,12 @@ def calculate_index_bounds(level_len, num_len, para_len, step_len, l, i, j, k): end_index = start_index + int(step_len) return start_index, end_index - def append_composite_coords(dates, tree_values, lat, coords): + def append_composite_coords(dates, tree_values, lat, coords, tags): for value in tree_values: coords[dates]["composite"].append([lat, value]) + coords[dates].setdefault("tags", []).append(tags) - def emit_leaf(lat, lon_values, result): + def emit_leaf(lat, lon_values, result, tags): lon_values = [float(val) for val in lon_values] if all(val is None for val in result): fields["dates"] = fields["dates"][:-1] @@ -559,7 +629,7 @@ def emit_leaf(lat, lon_values, result): para_len = num_len / len(fields["param"]) step_len = para_len / len(fields["step"]) - append_composite_coords(fields["dates"][-1], lon_values, lat, coords) + append_composite_coords(fields["dates"][-1], lon_values, lat, coords, tags) for l, level in enumerate(fields["levels"]): # noqa: E741 for i, num in enumerate(fields["number"]): @@ -577,7 +647,7 @@ def emit_leaf(lat, lon_values, result): for child in tree.children: # Compacted unstructured leaf: values=(lat, lon), own result. Emit directly. if is_merged_node(child): - emit_leaf(child.values[0], [child.values[1]], child.result) + emit_leaf(child.values[0], [child.values[1]], child.result, node_tags(child)) continue handle_non_leaf_node(child) result = handle_specific_axes(child) @@ -601,7 +671,7 @@ def emit_leaf(lat, lon_values, result): self.walk_tree_reforecast(child, fields, coords, mars_metadata, range_dict) else: - emit_leaf(fields["lat"], tree.values, tree.result) + emit_leaf(fields["lat"], tree.values, tree.result, node_tags(tree)) def walk_tree_step(self, tree, fields, coords, mars_metadata, range_dict): def create_composite_key_step(date, level, num, para): @@ -655,12 +725,13 @@ def calculate_index_bounds_step(level_len, num_len, para_len, step_len, l, i, j, end_index = start_index + int(step_len) return start_index, end_index - def append_composite_coords_step(dates, tree_values, lat, coords): + def append_composite_coords_step(dates, tree_values, lat, coords, tags): # for date in dates: for value in tree_values: coords[dates]["composite"].append([lat, value]) + coords[dates].setdefault("tags", []).append(tags) - def emit_leaf_step(lat, lon_values, result): + def emit_leaf_step(lat, lon_values, result, tags): """Emit one spatial leaf for the step walker (shared by legacy and merged).""" lon_values = [float(val) for val in lon_values] if all(val is None for val in result): @@ -680,7 +751,7 @@ def emit_leaf_step(lat, lon_values, result): para_len = level_len / len(fields["param"]) for date in fields["dates"]: - append_composite_coords_step(date, lon_values, lat, coords) + append_composite_coords_step(date, lon_values, lat, coords, tags) for d, date in enumerate(fields["dates"]): for l, level in enumerate(fields["levels"]): # noqa: E741 @@ -701,7 +772,7 @@ def emit_leaf_step(lat, lon_values, result): for child in tree.children: # Compacted unstructured leaf: values=(lat, lon), own result. Emit directly. if is_merged_node(child): - emit_leaf_step(child.values[0], [child.values[1]], child.result) + emit_leaf_step(child.values[0], [child.values[1]], child.result, node_tags(child)) continue handle_non_leaf_node_step(child) result = handle_specific_axes_step(child) @@ -727,7 +798,7 @@ def emit_leaf_step(lat, lon_values, result): self.walk_tree_step(child, fields, coords, mars_metadata, range_dict) else: - emit_leaf_step(fields["lat"], tree.values, tree.result) + emit_leaf_step(fields["lat"], tree.values, tree.result, node_tags(tree)) def walk_tree_month(self, tree, fields, coords, mars_metadata, range_dict, _ctx=None): """Walk the result tree for monthly-mean streams (e.g. clmn). @@ -776,11 +847,12 @@ def handle_specific_axes_month(child): return child.values return None - def append_composite_coords_month(date_key, tree_values, lat): + def append_composite_coords_month(date_key, tree_values, lat, tags): for value in tree_values: coords[date_key]["composite"].append([lat, value]) + coords[date_key].setdefault("tags", []).append(tags) - def emit_leaf_month(lat, lon_values, result): + def emit_leaf_month(lat, lon_values, result, tags): """Emit one spatial leaf for the month walker (shared by legacy and merged). ``lat`` is the latitude for this leaf, ``lon_values`` the leaf's list of @@ -840,7 +912,7 @@ def emit_leaf_month(lat, lon_values, result): # Append this leaf's longitude values to composite coords for # every date key in scope. for date in leaf_dates: - append_composite_coords_month(date, lon_values, lat) + append_composite_coords_month(date, lon_values, lat, tags) for d, date in enumerate(leaf_dates): for l, level in enumerate(fields["levels"]): # noqa: E741 @@ -857,7 +929,7 @@ def emit_leaf_month(lat, lon_values, result): for child in tree.children: # Compacted unstructured leaf: values=(lat, lon), own result. Emit directly. if is_merged_node(child): - emit_leaf_month(child.values[0], [child.values[1]], child.result) + emit_leaf_month(child.values[0], [child.values[1]], child.result, node_tags(child)) continue handle_non_leaf_node_month(child) result = handle_specific_axes_month(child) @@ -898,7 +970,7 @@ def emit_leaf_month(lat, lon_values, result): self.walk_tree_month(child, fields, coords, mars_metadata, range_dict, _ctx=child_ctx) else: - emit_leaf_month(fields["lat"], tree.values, tree.result) + emit_leaf_month(fields["lat"], tree.values, tree.result, node_tags(tree)) @abstractmethod def add_coverage(self, mars_metadata, coords, values): @@ -984,13 +1056,14 @@ def _reforecast_records(self, result): the full root-to-leaf path (latitude/longitude included), using the same ``itertools.product`` layout the compressed leaf ``result`` array follows. Handles both the compacted ``MergedTensorIndexNode`` lat/lon leaves and - the classic nested latitude/longitude branches. + the classic nested latitude/longitude branches. Each record also carries the + leaf's ``(index, label)`` tags under ``"__tags__"``. """ import itertools records = [] - def emit(full_path, flat_result): + def emit(full_path, flat_result, tags): axis_names = [name for name, _ in full_path] axis_values = [values for _, values in full_path] for idx, combo in enumerate(itertools.product(*axis_values)): @@ -999,11 +1072,12 @@ def emit(full_path, flat_result): continue records.append(dict(zip(axis_names, combo))) records[-1]["__value__"] = value + records[-1]["__tags__"] = tags def recurse(node, path): children = node.children if len(children) == 0: - emit(path, node.result) + emit(path, node.result, node_tags(node)) return for child in children: if is_merged_node(child): @@ -1011,6 +1085,7 @@ def recurse(node, path): emit( path + [("latitude", (lat,)), ("longitude", (lon,))], child.result, + node_tags(child), ) continue recurse(child, path + [(child.axis.name, tuple(child.values))]) @@ -1093,7 +1168,7 @@ def to_timedelta(value): if key not in coverages: meta = {} for name in rec: - if name == "__value__" or name in exclude_meta: + if name in ("__value__", "__tags__") or name in exclude_meta: continue meta[name] = stringify(rec[name]) meta["number"] = number diff --git a/tests/test_encoder_labels.py b/tests/test_encoder_labels.py new file mode 100644 index 0000000..cfee00e --- /dev/null +++ b/tests/test_encoder_labels.py @@ -0,0 +1,361 @@ +"""Tests for point labels carried as polytope tags. + +polytope-mars tags every requested point as ``(index, label)``: ``index`` is its +position in the request and ``label`` the user's label, or ``None``. Point-like +encoders emit one coverage per tag, in request order, so two requested points +that snap to the same grid point are both returned. ``mars:metadata.label`` is set +only when a label was requested. Trees without tags give the same output as before. +""" + +import logging + +import numpy as np +import pytest +from conftest import ( + MergedTensorIndexNode, + chain, + forecast_tree, + make_leaf, + make_merged_point, + make_point, + month_tree, + node, + reforecast_separate_datetime_tree, + reforecast_separate_datetime_vertical_tree, + tip, +) +from polytope_feature.datacube.tensor_index_tree import TensorIndexTree + +from covjsonkit.api import Covjsonkit +from covjsonkit.encoder.encoder import expand_points_by_tags, node_tags, parse_tag + + +def tagged(factory, tags_by_point): + """Wrap a conftest point factory so each point's leaf carries the given tags. + + The latitude node also gets the tags, as polytope stamps them on both. + """ + + def build(lat, lon, result): + point = factory(lat, lon, result) + tags = set(tags_by_point.get((lat, lon), ())) + point.tags.update(tags) + for child in point.children: + child.tags.update(tags) + return point + + return build + + +def encode(kind, tree, method="from_polytope"): + return getattr(Covjsonkit().encode("CoverageCollection", kind), method)(tree) + + +def labels(covjson): + return [c["mars:metadata"].get("label") for c in covjson["coverages"]] + + +def points_of(covjson): + return [ + (c["domain"]["axes"]["latitude"]["values"][0], c["domain"]["axes"]["longitude"]["values"][0]) + for c in covjson["coverages"] + ] + + +def values_of(covjson): + return [list(c["ranges"].values())[0]["values"] for c in covjson["coverages"]] + + +# Two grid points; results are per step (0, 6) +P1 = (48.0, 11.0, [1.0, 2.0]) +P2 = (50.0, 12.0, [3.0, 4.0]) + + +# -- Helpers -- + + +class TestTagHelpers: + def test_parse_tag(self): + assert parse_tag((0, "Lisbon")) == (0, "Lisbon") + assert parse_tag((3, None)) == (3, None) + # Anything that is not (int, label) is a bare label without an index + assert parse_tag("Lisbon") == (None, "Lisbon") + assert parse_tag((True, "x")) == (None, (True, "x")) + + def test_node_tags_sorted_by_index(self): + leaf = make_leaf(11.0, [1.0]) + leaf.tags.update({(2, "C"), (0, "A")}) + assert node_tags(leaf) == [(0, "A"), (2, "C")] + assert node_tags(make_leaf(11.0, [1.0])) == [] + + def test_expand_orders_by_index_and_duplicates(self): + tags = [[(2, "C")], [(0, "A"), (1, "B")]] + assert expand_points_by_tags(tags, 2) == [(1, "A"), (1, "B"), (0, "C")] + + def test_expand_without_tags_is_tree_order(self): + assert expand_points_by_tags(None, 3) == [(0, None), (1, None), (2, None)] + assert expand_points_by_tags([[], [], []], 3) == [(0, None), (1, None), (2, None)] + + def test_expand_untagged_after_tagged(self, caplog): + with caplog.at_level(logging.WARNING): + assert expand_points_by_tags([[], [(0, "A")]], 2) == [(1, "A"), (0, None)] + assert "carry no tag" in caplog.text + + def test_expand_mismatched_length_ignores_tags(self): + assert expand_points_by_tags([[(0, "A")]], 2) == [(0, None), (1, None)] + + +# -- TimeSeries (PointSeries) -- + + +class TestTimeSeriesLabels: + def test_labels_in_request_order(self): + # Tree order is P1, P2 but the request listed P2 first + tree = forecast_tree( + [P1, P2], + step=(0, 6), + point_factory=tagged(make_point, {P1[:2]: [(1, "Munich")], P2[:2]: [(0, "Bonn")]}), + ) + covjson = encode("pointseries", tree) + assert labels(covjson) == ["Bonn", "Munich"] + assert points_of(covjson) == [(50.0, 12.0), (48.0, 11.0)] + assert values_of(covjson) == [[3.0, 4.0], [1.0, 2.0]] + + def test_merged_points_duplicated(self): + # Two requested points snapped to the same grid point + tree = forecast_tree( + [P1], step=(0, 6), point_factory=tagged(make_point, {P1[:2]: [(0, "StationA"), (1, "StationB")]}) + ) + covjson = encode("pointseries", tree) + assert labels(covjson) == ["StationA", "StationB"] + a, b = covjson["coverages"] + assert a["domain"] == b["domain"] + assert a["ranges"] == b["ranges"] + + def test_unlabelled_same_output_as_labelled(self): + def build(label_a, label_b): + return forecast_tree( + [P1, P2], + step=(0, 6), + point_factory=tagged(make_point, {P1[:2]: [(0, label_a), (1, label_b)], P2[:2]: [(2, None)]}), + ) + + labelled = encode("pointseries", build("A", "B")) + unlabelled = encode("pointseries", build(None, None)) + assert labels(labelled) == ["A", "B", None] + assert labels(unlabelled) == [None, None, None] + assert all("label" not in c["mars:metadata"] for c in unlabelled["coverages"]) + for a, b in zip(labelled["coverages"], unlabelled["coverages"]): + a_meta = {k: v for k, v in a["mars:metadata"].items() if k != "label"} + assert a_meta == b["mars:metadata"] + assert a["domain"] == b["domain"] + assert a["ranges"] == b["ranges"] + + def test_untagged_tree_unchanged(self): + plain = encode("pointseries", forecast_tree([P1, P2], step=(0, 6))) + assert labels(plain) == [None, None] + assert points_of(plain) == [(48.0, 11.0), (50.0, 12.0)] + + def test_tags_taken_from_leaf_not_latitude_row(self): + # Two grid points in the same latitude row: the latitude node carries both + # tags, each longitude leaf only its own. + lat = node("latitude", (48.0,)) + lat.tags.update({(0, "West"), (1, "East")}) + west = make_leaf(11.0, [1.0, 2.0]) + west.tags.add((0, "West")) + east = make_leaf(12.0, [3.0, 4.0]) + east.tags.add((1, "East")) + lat.add_child(west) + lat.add_child(east) + tree = forecast_tree([], step=(0, 6)) + tip(tree).add_child(lat) + + covjson = encode("pointseries", tree) + assert labels(covjson) == ["West", "East"] + assert points_of(covjson) == [(48.0, 11.0), (48.0, 12.0)] + + def test_multiple_dates(self): + dates = (np.datetime64("2025-01-01T00:00:00"), np.datetime64("2025-01-02T00:00:00")) + tree = chain(TensorIndexTree(), node("class", ("od",))) + root = tip(tree) + factory = tagged(make_point, {P1[:2]: [(1, "B")], P2[:2]: [(0, "A")]}) + for date in dates: + branch = chain( + node("date", (date,)), + node("domain", ("g",)), + node("expver", ("0001",)), + node("levtype", ("sfc",)), + node("param", ("167",)), + node("step", (0, 6)), + node("stream", ("oper",)), + node("type", ("fc",)), + ) + parent = tip(branch) + parent.add_child(factory(*P1)) + parent.add_child(factory(*P2)) + root.add_child(branch) + + covjson = encode("pointseries", tree) + # Point-major in request order, then date + assert labels(covjson) == ["A", "A", "B", "B"] + assert [c["mars:metadata"]["Forecast date"] for c in covjson["coverages"]] == [ + "2025-01-01T00:00:00Z", + "2025-01-02T00:00:00Z", + ] * 2 + + def test_month(self): + tree = month_tree( + [(48.0, 11.0, [1.0, 2.0])], + years=(2020, 2021), + point_factory=tagged(make_point, {(48.0, 11.0): [(0, "A"), (1, "B")]}), + ) + covjson = encode("pointseries", tree, "from_polytope_month") + assert labels(covjson) == ["A", "B"] + assert values_of(covjson) == [[1.0, 2.0], [1.0, 2.0]] + + def test_reforecast_separate_datetime(self): + hdates = (np.datetime64("2025-07-14T00:00:00"), np.datetime64("2025-07-15T00:00:00")) + times = (np.timedelta64(0, "h"),) + tree = reforecast_separate_datetime_tree( + [(48.0, 11.0, [1.0, 2.0]), (50.0, 12.0, [3.0, 4.0])], + hdates, + times, + point_factory=tagged(make_point, {(48.0, 11.0): [(1, "B"), (2, "C")], (50.0, 12.0): [(0, "A")]}), + ) + covjson = encode("pointseries", tree, "from_polytope_reforecast") + assert labels(covjson) == ["A", "B", "C"] + assert values_of(covjson) == [[3.0, 4.0], [1.0, 2.0], [1.0, 2.0]] + + @pytest.mark.skipif(MergedTensorIndexNode is None, reason="polytope without MergedTensorIndexNode") + def test_merged_node_leaves(self): + tree = forecast_tree( + [P1, P2], + step=(0, 6), + point_factory=tagged(make_merged_point, {P1[:2]: [(1, "B")], P2[:2]: [(0, "A")]}), + ) + covjson = encode("pointseries", tree) + assert labels(covjson) == ["A", "B"] + assert points_of(covjson) == [(50.0, 12.0), (48.0, 11.0)] + + +def step_tree(factory): + """Tree with a separate ``time`` axis, for the ``from_polytope_step`` path.""" + tree = chain( + TensorIndexTree(), + node("class", ("od",)), + node("date", (np.datetime64("2025-01-01T00:00:00"),)), + node("domain", ("g",)), + node("expver", ("0001",)), + node("levtype", ("sfc",)), + node("param", ("167",)), + node("step", (0,)), + node("stream", ("oper",)), + node("time", (np.timedelta64(0, "h"), np.timedelta64(6, "h"))), + node("type", ("fc",)), + ) + parent = tip(tree) + parent.add_child(factory(*P1)) + parent.add_child(factory(*P2)) + return tree + + +@pytest.mark.parametrize("kind", ["pointseries", "position"]) +def test_step_path_labels(kind): + covjson = encode( + kind, step_tree(tagged(make_point, {P1[:2]: [(1, "B"), (2, "C")], P2[:2]: [(0, "A")]})), "from_polytope_step" + ) + assert labels(covjson) == ["A", "B", "C"] + assert values_of(covjson) == [[3.0, 4.0], [1.0, 2.0], [1.0, 2.0]] + + plain = encode(kind, step_tree(make_point), "from_polytope_step") + assert labels(plain) == [None, None] + assert values_of(plain) == [[1.0, 2.0], [3.0, 4.0]] + + +# -- Position -- + + +class TestPositionLabels: + def test_labels_and_duplicates(self): + tree = forecast_tree( + [P1, P2], + step=(0, 6), + point_factory=tagged(make_point, {P1[:2]: [(1, "B"), (2, "C")], P2[:2]: [(0, "A")]}), + ) + covjson = encode("position", tree) + assert labels(covjson) == ["A", "B", "C"] + assert points_of(covjson) == [(50.0, 12.0), (48.0, 11.0), (48.0, 11.0)] + assert values_of(covjson)[1] == values_of(covjson)[2] + + def test_untagged_unchanged(self): + covjson = encode("position", forecast_tree([P1, P2], step=(0, 6))) + assert labels(covjson) == [None, None] + assert points_of(covjson) == [(48.0, 11.0), (50.0, 12.0)] + + def test_reforecast_separate_datetime(self): + hdates = (np.datetime64("2025-07-14T00:00:00"),) + times = (np.timedelta64(0, "h"), np.timedelta64(12, "h")) + tree = reforecast_separate_datetime_tree( + [(48.0, 11.0, [1.0, 2.0]), (50.0, 12.0, [3.0, 4.0])], + hdates, + times, + point_factory=tagged(make_point, {(48.0, 11.0): [(1, "B")], (50.0, 12.0): [(0, "A")]}), + ) + covjson = encode("position", tree, "from_polytope_reforecast") + # One coverage per (point, reference); request order first + assert labels(covjson) == ["A", "A", "B", "B"] + assert all("__tags__" not in c["mars:metadata"] for c in covjson["coverages"]) + + +# -- VerticalProfile -- + + +def vp_tree(points, factory): + tree = chain( + TensorIndexTree(), + node("class", ("od",)), + node("date", (np.datetime64("2025-01-01T00:00:00"),)), + node("domain", ("g",)), + node("expver", ("0001",)), + node("levtype", ("pl",)), + node("param", ("130",)), + node("step", (0,)), + node("stream", ("oper",)), + node("type", ("an",)), + node("levelist", (1000, 850)), + ) + parent = tip(tree) + for lat, lon, result in points: + parent.add_child(factory(lat, lon, result)) + return tree + + +class TestVerticalProfileLabels: + def test_labels_and_duplicates(self): + tree = vp_tree( + [(48.0, 11.0, [290.0, 280.0]), (50.0, 12.0, [291.0, 281.0])], + tagged(make_point, {(48.0, 11.0): [(1, "B"), (2, "C")], (50.0, 12.0): [(0, "A")]}), + ) + covjson = encode("verticalprofile", tree) + assert labels(covjson) == ["A", "B", "C"] + assert values_of(covjson) == [[291.0, 281.0], [290.0, 280.0], [290.0, 280.0]] + + def test_untagged_unchanged(self): + tree = vp_tree([(48.0, 11.0, [290.0, 280.0]), (50.0, 12.0, [291.0, 281.0])], make_point) + covjson = encode("verticalprofile", tree) + assert labels(covjson) == [None, None] + assert values_of(covjson) == [[290.0, 280.0], [291.0, 281.0]] + + def test_reforecast_separate_datetime(self): + hdates = (np.datetime64("2025-07-14T00:00:00"),) + times = (np.timedelta64(0, "h"),) + tree = reforecast_separate_datetime_vertical_tree( + [(48.0, 11.0, [1.0, 2.0])], + hdates, + times, + (500, 850), + point_factory=tagged(make_point, {(48.0, 11.0): [(0, "A"), (1, "B")]}), + ) + covjson = encode("verticalprofile", tree, "from_polytope_reforecast") + assert labels(covjson) == ["A", "B"] + assert values_of(covjson) == [[1.0, 2.0], [1.0, 2.0]] From 948c1ce61720d90bb347fc6b606b4101f49068f0 Mon Sep 17 00:00:00 2001 From: awarde96 Date: Thu, 1 Oct 2026 07:58:21 +0000 Subject: [PATCH 5/5] Align decoders with label changes --- covjsonkit/decoder/Position.py | 62 +++++----- covjsonkit/decoder/TimeSeries.py | 61 +++++----- covjsonkit/decoder/VerticalProfile.py | 50 ++++---- covjsonkit/decoder/decoder.py | 33 ++++++ tests/test_encoder_labels.py | 158 ++++++++++++++++++++++---- 5 files changed, 253 insertions(+), 111 deletions(-) diff --git a/covjsonkit/decoder/Position.py b/covjsonkit/decoder/Position.py index 4331582..c6d426c 100644 --- a/covjsonkit/decoder/Position.py +++ b/covjsonkit/decoder/Position.py @@ -109,25 +109,29 @@ def to_xarray(self): dims = ["latitude", "longitude", "levelist", "number", "datetime", "t"] ds = [] - unique_coords = set() # To track unique coordinate tuples - unique_domains = [] # To store unique domains - - for domain in self.domains: - # Extract coordinate values - x = domain["axes"][self.x_name]["values"][0] - y = domain["axes"][self.y_name]["values"][0] - z = domain["axes"][self.z_name]["values"][0] - t = tuple(domain["axes"]["t"]["values"]) # Use tuple for hashable type - - # Create a unique identifier for the domain - coord_tuple = (x, y, z, t) - - # Check if this coordinate combination is already seen - if coord_tuple not in unique_coords: - unique_coords.add(coord_tuple) # Mark as seen - unique_domains.append(domain) # Add to unique domains - - all_coords = unique_domains + # One dataset per requested point: coverages are grouped by domain and label, + # and points that snapped to the same grid point are kept apart. + def point_key(coverage): + axes = coverage["domain"]["axes"] + return ( + axes[self.x_name]["values"][0], + axes[self.y_name]["values"][0], + axes[self.z_name]["values"][0], + ) + + def slot_key(coverage): + return (coverage["mars:metadata"]["number"], coverage["mars:metadata"]["Forecast date"]) + + # Within each point, one dataset per distinct time axis (as before), filled from + # that point's coverages only. + datasets = [] + for group in self._point_groups(point_key, slot_key): + seen_t = set() + for coverage in group: + t = tuple(coverage["domain"]["axes"]["t"]["values"]) + if t not in seen_t: + seen_t.add(t) + datasets.append((group, coverage["domain"])) num = [] datetime = [] @@ -137,8 +141,8 @@ def to_xarray(self): nums = list(set(num)) datetime = list(set(datetime)) - # Process each coordinate domain - for coords in all_coords: + # Process each requested point and time axis + for group, coords in datasets: dataarraydict = {} x = coords["axes"][self.x_name]["values"] y = coords["axes"][self.y_name]["values"] @@ -147,7 +151,7 @@ def to_xarray(self): steps = [step.replace("Z", "") for step in steps] steps = pd.to_datetime(steps) - cov_idx_list = self._find_coverages(nums, datetime, x, y, z) + cov_idx_list = self._find_coverages(nums, datetime, x, y, z, group) coords = { "latitude": x, @@ -181,25 +185,21 @@ def to_xarray(self): attrs, ) - ds.append(xr.Dataset(data_vars=dataarraydict, coords=coords)) - - # Combine all DataArrays into a Dataset - for mars_metadata in self.mars_metadata[0]: - if mars_metadata != "date" and mars_metadata != "step": - for dss in ds: - dss.attrs[mars_metadata] = self.mars_metadata[0][mars_metadata] + dss = xr.Dataset(data_vars=dataarraydict, coords=coords) + dss.attrs.update(self._point_dataset_attrs(group)) + ds.append(dss) if len(ds) == 1: return ds[0] return ds - def _find_coverages(self, nums, datetime, x, y, z): + def _find_coverages(self, nums, datetime, x, y, z, coverages=None): """Find coverages matching domain parameters and return with indices.""" result = [] for i, num in enumerate(nums): for j, date in enumerate(datetime): - for coverage in self.covjson["coverages"]: + for coverage in coverages if coverages is not None else self.covjson["coverages"]: if self._covers_domain(coverage, num, date, x, y, z): result.append((i, j, coverage)) return result diff --git a/covjsonkit/decoder/TimeSeries.py b/covjsonkit/decoder/TimeSeries.py index 28f5702..64a81b9 100644 --- a/covjsonkit/decoder/TimeSeries.py +++ b/covjsonkit/decoder/TimeSeries.py @@ -117,28 +117,29 @@ def to_xarray(self): dims = ["latitude", "longitude", "levelist", "number", "datetime", "t"] ds = [] - # Get coordinates for all domains - all_coords = self.get_domains() - - unique_coords = set() # To track unique coordinate tuples - unique_domains = [] # To store unique domains - - for domain in self.domains: - # Extract coordinate values - x = domain["axes"][self.x_name]["values"][0] - y = domain["axes"][self.y_name]["values"][0] - z = domain["axes"][self.z_name]["values"][0] - t = tuple(domain["axes"]["t"]["values"]) # Use tuple for hashable type - - # Create a unique identifier for the domain - coord_tuple = (x, y, z, t) + # One dataset per requested point: coverages are grouped by domain and label, + # and points that snapped to the same grid point are kept apart. + def point_key(coverage): + axes = coverage["domain"]["axes"] + return ( + axes[self.x_name]["values"][0], + axes[self.y_name]["values"][0], + axes[self.z_name]["values"][0], + ) - # Check if this coordinate combination is already seen - if coord_tuple not in unique_coords: - unique_coords.add(coord_tuple) # Mark as seen - unique_domains.append(domain) # Add to unique domains + def slot_key(coverage): + return (coverage["mars:metadata"]["number"], coverage["mars:metadata"]["Forecast date"]) - all_coords = unique_domains + # Within each point, one dataset per distinct time axis (as before), filled from + # that point's coverages only. + datasets = [] + for group in self._point_groups(point_key, slot_key): + seen_t = set() + for coverage in group: + t = tuple(coverage["domain"]["axes"]["t"]["values"]) + if t not in seen_t: + seen_t.add(t) + datasets.append((group, coverage["domain"])) num = [] datetime = [] @@ -148,8 +149,8 @@ def to_xarray(self): nums = list(set(num)) datetime = list(set(datetime)) - # Process each coordinate domain - for coords in all_coords: + # Process each requested point and time axis + for group, coords in datasets: dataarraydict = {} x = coords["axes"][self.x_name]["values"] y = coords["axes"][self.y_name]["values"] @@ -158,7 +159,7 @@ def to_xarray(self): steps = [step.replace("Z", "") for step in steps] steps = pd.to_datetime(steps) - cov_idx_list = self._find_coverages(nums, datetime, x, y, z) + cov_idx_list = self._find_coverages(nums, datetime, x, y, z, group) coords = { "latitude": x, @@ -192,26 +193,22 @@ def to_xarray(self): attrs, ) - ds.append(xr.Dataset(data_vars=dataarraydict, coords=coords)) - - # Combine all DataArrays into a Dataset - for mars_metadata in self.mars_metadata[0]: - if mars_metadata != "date" and mars_metadata != "step": - for dss in ds: - dss.attrs[mars_metadata] = self.mars_metadata[0][mars_metadata] + dss = xr.Dataset(data_vars=dataarraydict, coords=coords) + dss.attrs.update(self._point_dataset_attrs(group)) + ds.append(dss) if len(ds) == 1: return ds[0] return ds - def _find_coverages(self, nums, datetime, x, y, z): + def _find_coverages(self, nums, datetime, x, y, z, coverages=None): """Find coverages that match the given domain parameters (num, date, x, y, z) and return them along with domain parameter indices.""" result = [] for i, num in enumerate(nums): for j, date in enumerate(datetime): - for coverage in self.covjson["coverages"]: + for coverage in coverages if coverages is not None else self.covjson["coverages"]: if self._covers_domain(coverage, num, date, x, y, z): result.append((i, j, coverage)) return result diff --git a/covjsonkit/decoder/VerticalProfile.py b/covjsonkit/decoder/VerticalProfile.py index 23dc1f4..cb0e620 100644 --- a/covjsonkit/decoder/VerticalProfile.py +++ b/covjsonkit/decoder/VerticalProfile.py @@ -112,34 +112,29 @@ def to_xarray(self): ] ds = [] - # Get coordinates for all domains - all_coords = self.get_domains() - - unique_coords = set() # To track unique coordinate tuples - unique_domains = [] # To store unique domains - - for domain in self.domains: - # Extract coordinate values - x = domain["axes"][self.x_name]["values"][0] - y = domain["axes"][self.y_name]["values"][0] - z = domain["axes"][self.z_name]["values"][0] - - # Create a unique identifier for the domain - coord_tuple = (x, y, z) - - # Check if this coordinate combination is already seen - if coord_tuple not in unique_coords: - unique_coords.add(coord_tuple) # Mark as seen - unique_domains.append(domain) # Add to unique domains - - all_coords = unique_domains + # One dataset per requested point: coverages are grouped by domain and label, + # and points that snapped to the same grid point are kept apart. + def domain_key(coverage): + axes = coverage["domain"]["axes"] + return ( + axes[self.x_name]["values"][0], + axes[self.y_name]["values"][0], + axes[self.z_name]["values"][0], + ) + + def slot_key(coverage): + meta = coverage["mars:metadata"] + return (meta["number"], meta["Forecast date"], meta["step"]) + + groups = self._point_groups(domain_key, slot_key) param_values = {} # Initialize parameter values for all parameters for parameter in self.parameters: param_values[parameter] = [] - for domain_idx, coords in enumerate(all_coords): + for domain_idx, group in enumerate(groups): + coords = group[0]["domain"] dataarraydict = {} # Get coordinates @@ -175,7 +170,7 @@ def to_xarray(self): for k, step in enumerate(steps): if len(param_values[parameter][domain_idx][i][j]) <= k: param_values[parameter][domain_idx][i][j].append([]) - for coverage in self.covjson["coverages"]: + for coverage in group: new_step = ( dt.fromisoformat(date.replace("Z", "")) + timedelta(hours=parse_step_string(step)) ).isoformat() + "Z" @@ -211,12 +206,9 @@ def to_xarray(self): dataarray.attrs["long_name"] = self.get_parameter_metadata(parameter)["observedProperty"]["id"] dataarraydict[dataarray.attrs["long_name"]] = dataarray - ds.append(xr.Dataset(dataarraydict)) - - for mars_metadata in self.mars_metadata[0]: - for dss in ds: - if mars_metadata != "date" and mars_metadata != "step": - dss.attrs[mars_metadata] = self.mars_metadata[0][mars_metadata] + dss = xr.Dataset(dataarraydict) + dss.attrs.update(self._point_dataset_attrs(group)) + ds.append(dss) if len(ds) == 1: return ds[0] diff --git a/covjsonkit/decoder/decoder.py b/covjsonkit/decoder/decoder.py index 6922e46..4ba12c0 100644 --- a/covjsonkit/decoder/decoder.py +++ b/covjsonkit/decoder/decoder.py @@ -91,6 +91,39 @@ def get_mars_metadata(self): mars_metadata.append(coverage["mars:metadata"]) return mars_metadata + def _point_groups(self, domain_key, slot_key): + """Group coverages into one group per requested point, in first-seen order. + + Coverages with the same ``domain_key`` (spatial/time domain) and ``label`` + belong to the same point. When several requested points snapped to the same + grid point, their coverages repeat for each ``slot_key`` (e.g. number and + forecast date), so the n-th coverage for a given domain, label and slot + belongs to the n-th such point. + """ + groups = {} + order = [] + seen = {} + for coverage in self.covjson["coverages"]: + label = coverage.get("mars:metadata", {}).get("label") + domain = domain_key(coverage) + slot = (domain, label, slot_key(coverage)) + occurrence = seen.get(slot, 0) + seen[slot] = occurrence + 1 + key = (domain, label, occurrence) + if key not in groups: + groups[key] = [] + order.append(key) + groups[key].append(coverage) + return [groups[key] for key in order] + + def _point_dataset_attrs(self, group): + """Dataset attributes for one point: shared MARS metadata plus the point's own label.""" + attrs = {key: val for key, val in self.mars_metadata[0].items() if key not in ("date", "step", "label")} + label = group[0].get("mars:metadata", {}).get("label") + if label is not None: + attrs["label"] = label + return attrs + @abstractmethod def get_ranges(self): pass diff --git a/tests/test_encoder_labels.py b/tests/test_encoder_labels.py index cfee00e..137e837 100644 --- a/tests/test_encoder_labels.py +++ b/tests/test_encoder_labels.py @@ -71,6 +71,29 @@ def values_of(covjson): P2 = (50.0, 12.0, [3.0, 4.0]) +def two_date_tree(factory): + """Two forecast dates, each with points P1 and P2.""" + dates = (np.datetime64("2025-01-01T00:00:00"), np.datetime64("2025-01-02T00:00:00")) + tree = chain(TensorIndexTree(), node("class", ("od",))) + root = tip(tree) + for date in dates: + branch = chain( + node("date", (date,)), + node("domain", ("g",)), + node("expver", ("0001",)), + node("levtype", ("sfc",)), + node("param", ("167",)), + node("step", (0, 6)), + node("stream", ("oper",)), + node("type", ("fc",)), + ) + parent = tip(branch) + parent.add_child(factory(*P1)) + parent.add_child(factory(*P2)) + root.add_child(branch) + return tree + + # -- Helpers -- @@ -175,25 +198,7 @@ def test_tags_taken_from_leaf_not_latitude_row(self): assert points_of(covjson) == [(48.0, 11.0), (48.0, 12.0)] def test_multiple_dates(self): - dates = (np.datetime64("2025-01-01T00:00:00"), np.datetime64("2025-01-02T00:00:00")) - tree = chain(TensorIndexTree(), node("class", ("od",))) - root = tip(tree) - factory = tagged(make_point, {P1[:2]: [(1, "B")], P2[:2]: [(0, "A")]}) - for date in dates: - branch = chain( - node("date", (date,)), - node("domain", ("g",)), - node("expver", ("0001",)), - node("levtype", ("sfc",)), - node("param", ("167",)), - node("step", (0, 6)), - node("stream", ("oper",)), - node("type", ("fc",)), - ) - parent = tip(branch) - parent.add_child(factory(*P1)) - parent.add_child(factory(*P2)) - root.add_child(branch) + tree = two_date_tree(tagged(make_point, {P1[:2]: [(1, "B")], P2[:2]: [(0, "A")]})) covjson = encode("pointseries", tree) # Point-major in request order, then date @@ -359,3 +364,118 @@ def test_reforecast_separate_datetime(self): covjson = encode("verticalprofile", tree, "from_polytope_reforecast") assert labels(covjson) == ["A", "B"] assert values_of(covjson) == [[1.0, 2.0], [1.0, 2.0]] + + +# -- Decoding (to_xarray) -- + + +def decoded(covjson): + ds = Covjsonkit().decode(covjson).to_xarray() + return ds if isinstance(ds, list) else [ds] + + +def summary(datasets): + out = [] + for ds in datasets: + var = list(ds.data_vars)[0] + out.append( + ( + float(ds["latitude"].values[0]), + float(ds["longitude"].values[0]), + ds.attrs.get("label"), + np.asarray(ds[var].values).ravel().tolist(), + ) + ) + return out + + +def numbered_tree(factory, numbers=(1, 2)): + """Forecast tree with several ensemble members; results are [step, number] flattened per point.""" + tree = chain( + TensorIndexTree(), + node("class", ("od",)), + node("date", (np.datetime64("2025-01-01T00:00:00"),)), + node("domain", ("g",)), + node("expver", ("0001",)), + node("levtype", ("sfc",)), + node("number", numbers), + node("param", ("167",)), + node("step", (0, 6)), + node("stream", ("enfo",)), + node("type", ("pf",)), + ) + parent = tip(tree) + parent.add_child(factory(48.0, 11.0, [1.0, 2.0, 10.0, 20.0])) + parent.add_child(factory(50.0, 12.0, [3.0, 4.0, 30.0, 40.0])) + return tree + + +@pytest.mark.parametrize("kind", ["pointseries", "position"]) +class TestPointDecodersWithLabels: + def test_one_dataset_per_label(self, kind): + tree = forecast_tree( + [P1, P2], step=(0, 6), point_factory=tagged(make_point, {P1[:2]: [(1, "B")], P2[:2]: [(0, "A")]}) + ) + assert summary(decoded(encode(kind, tree))) == [ + (50.0, 12.0, "A", [3.0, 4.0]), + (48.0, 11.0, "B", [1.0, 2.0]), + ] + + def test_merged_points_kept_apart(self, kind): + tags = {P1[:2]: [(1, "B"), (2, "C")], P2[:2]: [(0, "A")]} + tree = forecast_tree([P1, P2], step=(0, 6), point_factory=tagged(make_point, tags)) + assert summary(decoded(encode(kind, tree))) == [ + (50.0, 12.0, "A", [3.0, 4.0]), + (48.0, 11.0, "B", [1.0, 2.0]), + (48.0, 11.0, "C", [1.0, 2.0]), + ] + + def test_merged_unlabelled_kept_apart_with_numbers(self, kind): + tags = {(48.0, 11.0): [(1, None), (2, None)], (50.0, 12.0): [(0, None)]} + datasets = decoded(encode(kind, numbered_tree(tagged(make_point, tags)))) + assert len(datasets) == 3 + assert [(lat, lon, label) for lat, lon, label, _ in summary(datasets)] == [ + (50.0, 12.0, None), + (48.0, 11.0, None), + (48.0, 11.0, None), + ] + # Each dataset holds both ensemble members + assert all(ds.sizes["number"] == 2 for ds in datasets) + assert datasets[1].identical(datasets[2]) + + def test_untagged_unchanged(self, kind): + datasets = decoded(encode(kind, numbered_tree(make_point))) + assert [(lat, lon, label) for lat, lon, label, _ in summary(datasets)] == [ + (48.0, 11.0, None), + (50.0, 12.0, None), + ] + assert all("label" not in ds.attrs for ds in datasets) + + +class TestVerticalProfileDecoderWithLabels: + def test_merged_points_kept_apart(self): + tree = vp_tree( + [(48.0, 11.0, [290.0, 280.0]), (50.0, 12.0, [291.0, 281.0])], + tagged(make_point, {(48.0, 11.0): [(1, "B"), (2, "C")], (50.0, 12.0): [(0, "A")]}), + ) + assert summary(decoded(encode("verticalprofile", tree))) == [ + (50.0, 12.0, "A", [291.0, 281.0]), + (48.0, 11.0, "B", [290.0, 280.0]), + (48.0, 11.0, "C", [290.0, 280.0]), + ] + + def test_untagged_unchanged(self): + tree = vp_tree([(48.0, 11.0, [290.0, 280.0]), (50.0, 12.0, [291.0, 281.0])], make_point) + assert summary(decoded(encode("verticalprofile", tree))) == [ + (48.0, 11.0, None, [290.0, 280.0]), + (50.0, 12.0, None, [291.0, 281.0]), + ] + + +@pytest.mark.parametrize("kind", ["pointseries", "position"]) +def test_decode_multiple_dates_with_merged_labels(kind): + tags = {P1[:2]: [(1, "B"), (2, "C")], P2[:2]: [(0, "A")]} + datasets = decoded(encode(kind, two_date_tree(tagged(make_point, tags)))) + # As before: one dataset per point per distinct time axis, each holding both dates + assert [ds.attrs.get("label") for ds in datasets] == ["A", "A", "B", "B", "C", "C"] + assert all(ds.sizes["datetime"] == 2 for ds in datasets)