"""Subset extraction helpers for ARL meteorology files."""
from __future__ import annotations
import os
from collections import OrderedDict
from collections.abc import Iterable, Sequence
from pathlib import Path
from typing import TYPE_CHECKING, BinaryIO
from arlmet._io import atomic_output, reject_same_file
from arlmet.exceptions import ARLFormatError
from arlmet.file import File
from arlmet.grid import Grid, GridWindow
from arlmet.header import Header, record_length_from_grid, split_grid_component
from arlmet.index import IndexRecord, LvlInfo, VarInfo, _derive_index_forecast
from arlmet.packing import calculate_checksum
from arlmet.record import DataRecord
from arlmet.vertical import VerticalAxis
if TYPE_CHECKING:
from arlmet.recordset import RecordSet
def normalize_levels(
vertical_axis: VerticalAxis, levels: Iterable[int] | None
) -> tuple[int, ...]:
"""
Normalize a level selection to sorted unique ARL level indices.
"""
if levels is None:
return tuple(range(len(vertical_axis.levels)))
normalized = tuple(sorted({int(level) for level in levels}))
if not normalized:
raise ValueError("levels must include at least one level index.")
max_index = len(vertical_axis.levels) - 1
if normalized[0] < 0 or normalized[-1] > max_index:
raise ValueError(f"levels must be between 0 and {max_index}, got {normalized}.")
return normalized
def resolve_window(
file: File, bbox: tuple[float, float, float, float] | None
) -> GridWindow:
"""
Resolve a bbox selection to a grid window.
"""
if bbox is None:
return file.grid.full_window()
return file.grid.window_from_bbox(bbox)
def select_records(
records: Sequence[DataRecord],
*,
levels: set[int] | None = None,
variables: set[str] | None = None,
) -> list[DataRecord]:
"""
Filter records by ARL level index and variable name.
"""
return [
record
for record in records
if (levels is None or record.level in levels)
and (variables is None or record.variable in variables)
]
def _source_index_entries(
selected_records: Sequence[DataRecord], level_map: dict[int, int]
) -> list[tuple[int, str, VarInfo]]:
"""
Index manifest entries ``(output level, name, info)`` for selected records.
Each record is followed by its DIF record, if any, as in the output file.
Checksums and reserved flags come from the source index; they size the
output index record for validation.
"""
entries = []
for record in selected_records:
level = level_map[record.level]
for rec in (record,) if record.diff is None else (record, record.diff):
entries.append(
(
level,
rec.variable,
VarInfo(checksum=rec.checksum, reserved=(rec._reserved or "")[:1]),
)
)
return entries
def _build_subset_index_record(
recordset: RecordSet,
*,
subset_grid: Grid,
subset_axis: VerticalAxis,
selected_records: Sequence[DataRecord],
entries: Iterable[tuple[int, str, VarInfo]],
) -> IndexRecord:
"""
Build the output index record for one subsetted time step.
``entries`` lists ``(output level, variable name, VarInfo)`` in file order.
"""
forecast = _derive_index_forecast(
(record.forecast for record in selected_records),
recordset.forecast,
)
level_records: dict[int, OrderedDict[str, VarInfo]] = {
level: OrderedDict() for level in range(len(subset_axis.levels))
}
for level, name, info in entries:
level_records[level][name] = info
grid_x = split_grid_component(subset_grid.nx)[0]
grid_y = split_grid_component(subset_grid.ny)[0]
levels = [
LvlInfo(
level=level,
height=float(height),
variables=level_records[level],
)
for level, height in enumerate(subset_axis.levels)
]
projection = subset_grid.projection
time = recordset.time
return IndexRecord(
header=Header(
year=time.year,
month=time.month,
day=time.day,
hour=time.hour,
forecast=forecast,
level=0,
grid=(grid_x, grid_y),
variable="INDX",
exponent=0,
precision=0.0,
initial_value=0.0,
),
source=recordset.source,
forecast=forecast,
minutes=time.minute,
pole_lat=projection.pole_lat,
pole_lon=projection.pole_lon,
tangent_lat=projection.tangent_lat,
tangent_lon=projection.tangent_lon,
grid_size=projection.grid_size,
orientation=projection.orientation,
cone_angle=projection.cone_angle,
sync_x=projection.sync_x,
sync_y=projection.sync_y,
sync_lat=projection.sync_lat,
sync_lon=projection.sync_lon,
reserved=subset_axis.offset,
nx=subset_grid.nx,
ny=subset_grid.ny,
nz=len(levels),
vertical_flag=subset_axis.flag,
levels=levels,
)
def validate_subset_record_length(
selected_recordsets: Sequence[tuple[RecordSet, Sequence[DataRecord]]],
*,
subset_grid: Grid,
subset_axis: VerticalAxis,
level_map: dict[int, int],
) -> None:
"""
Fail early when a cropped ARL grid cannot fit its index record.
"""
record_len = record_length_from_grid(grid=subset_grid)
for recordset, selected_records in selected_recordsets:
index = _build_subset_index_record(
recordset,
subset_grid=subset_grid,
subset_axis=subset_axis,
selected_records=selected_records,
entries=_source_index_entries(selected_records, level_map),
)
index_len = len(index.tobytes())
if index_len > record_len:
min_cells = index_len - Header.N_BYTES
raise ValueError(
"Subset grid is too small to encode an ARL index record: "
f"time {recordset.time} needs {index_len} bytes, but each record is "
f"only {record_len} bytes for grid {subset_grid.nx}x{subset_grid.ny}. "
f"The bbox must yield at least {min_cells} grid cells (nx*ny). "
"Expand the bbox or reduce levels/variables."
)
# Byte offsets of the level field in a record header (see Header.FIELDS).
_LEVEL_START, _LEVEL_STOP = Header.FIELDS["level"][:2]
def _copy_subset_records(
src: File,
out: BinaryIO,
selected_recordsets: Sequence[tuple[RecordSet, Sequence[DataRecord]]],
*,
subset_grid: Grid,
subset_axis: VerticalAxis,
level_map: dict[int, int],
) -> None:
"""
Write an uncropped subset by copying record bytes instead of re-packing.
Without a horizontal crop a record's packed bytes do not depend on which
other records are kept, so each selected record (and its DIF record) is
copied verbatim with only the header's level field renumbered. Only the
index records are rebuilt, with checksums recomputed from the copied
bytes and blank reserved flags, as a re-pack writes them.
The output decodes to exactly the source values. For records that
re-packing reproduces byte for byte (those written by arlmet without a
DIF record) it is identical to the re-pack path's output; otherwise the
re-pack path re-quantizes the values and this path is the more faithful.
Each time step's index record is written last, over a placeholder, once
the checksums of its records are known, so memory stays at a few records.
"""
record_length = src.record_length
src_fh = src.handle
placeholder = bytes(record_length)
for recordset, selected_records in selected_recordsets:
index_position = out.tell()
out.write(placeholder)
entries: list[tuple[int, str, VarInfo]] = []
for record in selected_records:
level = level_map[record.level]
level_field = f"{level:2d}".encode("ascii")
for rec in (record,) if record.diff is None else (record, record.diff):
src_fh.seek(rec.position)
raw = src_fh.read(record_length)
if len(raw) != record_length:
raise ARLFormatError(
f"{src.path}: record {rec.variable!r} at byte "
f"{rec.position} is truncated."
)
header = Header.from_bytes(raw[: Header.N_BYTES])
if header.variable != rec.variable or header.level != rec.level:
raise ARLFormatError(
f"DataRecord header mismatch at position {rec.position}: "
f"expected variable '{rec.variable}' level {rec.level}, "
f"got variable '{header.variable}' level '{header.level}'"
)
# Write the record with its level renumbered, without copying it.
view = memoryview(raw)
out.write(view[:_LEVEL_START])
out.write(level_field)
out.write(view[_LEVEL_STOP:])
checksum = calculate_checksum(view[Header.N_BYTES :])
entries.append((level, rec.variable, VarInfo(checksum, reserved="")))
index = _build_subset_index_record(
recordset,
subset_grid=subset_grid,
subset_axis=subset_axis,
selected_records=selected_records,
entries=entries,
)
out.seek(index_position)
out.write(index.to_record_bytes(record_length))
out.seek(0, os.SEEK_END)
def _repack_subset_records(
src: File,
dest: Path,
selected_recordsets: Sequence[tuple[RecordSet, Sequence[DataRecord]]],
*,
window: GridWindow,
subset_grid: Grid,
subset_axis: VerticalAxis,
level_map: dict[int, int],
) -> None:
"""
Write a cropped subset by unpacking each record's window and re-packing it.
"""
with File(
dest,
mode="w",
source=src.source,
grid=subset_grid,
vertical_axis=subset_axis,
) as out:
for src_recordset, selected_records in selected_recordsets:
dst_recordset = out.create_recordset(
src_recordset.time,
forecast=src_recordset.forecast,
)
for record in selected_records:
# record.read() returns the full-precision value
# (parent + diff when a diff is attached); the diff branch
# below relies on this so that create_datarecord(diff=...)
# can recompute the diff against the newly packed parent.
data = record.read(window=window)
dst_recordset.create_datarecord(
variable=record.variable,
level=level_map[record.level],
forecast=record.forecast,
data=data,
diff=record.diff.variable if record.diff is not None else None,
)
# Write each time step as soon as it is filled so peak memory
# is one time step, not the whole output.
out.flush()
def _is_full_window(grid: Grid, window: GridWindow) -> bool:
"""Return True when ``window`` covers all of ``grid`` (no horizontal crop)."""
return window == grid.full_window()