diff --git a/src/itzi_core/data_containers.py b/src/itzi_core/data_containers.py index 56045b4..2d11405 100644 --- a/src/itzi_core/data_containers.py +++ b/src/itzi_core/data_containers.py @@ -19,7 +19,15 @@ from typing import TYPE_CHECKING import numpy as np -from pydantic import BaseModel, ConfigDict, Field, NonNegativeFloat, NonNegativeInt, PositiveFloat +from pydantic import ( + BaseModel, + ConfigDict, + Field, + NonNegativeFloat, + NonNegativeInt, + PositiveFloat, + field_validator, +) from itzi_core.const import DefaultValues, InfiltrationModelType, TemporalType from itzi_core.providers.domain_data import DomainData @@ -146,7 +154,7 @@ class SimulationData(BaseModel): sim_time: datetime time_step: float # time step duration time_steps_counter: int # number of time steps since last update - continuity_data: ContinuityData | None # Made optional for use in tests + continuity_data: ContinuityData raw_arrays: dict[str, np.ndarray] accumulation_arrays: dict[str, np.ndarray] cell_dx: PositiveFloat # cell size in east-west direction @@ -211,8 +219,8 @@ class SimulationConfig(BaseModel): # Hotstart config hotstart_config: HotstartRunConfig | None = None # Input and output raster maps - input_map_names: dict[str, str | None] - output_map_names: dict[str, str | None] + input_map_names: dict[str, str] + output_map_names: dict[str, str] # Surface flow parameters surface_flow_parameters: SurfaceFlowParameters # Hydrology parameters @@ -225,6 +233,14 @@ class SimulationConfig(BaseModel): free_weir_coeff: NonNegativeFloat = Field(DefaultValues.FREE_WEIR_COEFF, ge=0, le=1) submerged_weir_coeff: NonNegativeFloat = Field(DefaultValues.SUBMERGED_WEIR_COEFF, ge=0, le=1) + @field_validator("input_map_names", "output_map_names", mode="before") + @classmethod + def remove_inactive_map_names(cls, value: object) -> object: + """Normalize legacy null-valued map entries to omitted inactive entries.""" + if isinstance(value, dict): + return {key: map_name for key, map_name in value.items() if map_name is not None} + return value + def as_str_dict(self) -> dict: """Convert the configuration to a dictionary with string representations.""" raw_dict = self.model_dump() diff --git a/src/itzi_core/drainage.py b/src/itzi_core/drainage.py index 4124501..d4e9947 100644 --- a/src/itzi_core/drainage.py +++ b/src/itzi_core/drainage.py @@ -14,27 +14,29 @@ from __future__ import annotations -from typing import TYPE_CHECKING import math +import tempfile from datetime import timedelta from enum import StrEnum from io import BytesIO -import tempfile +from typing import TYPE_CHECKING -import pyswmm import numpy as np +import pyswmm +from pyswmm.toolkitapi import NodeResults, SimulationParameters, SimulationTime from itzi_core import DefaultValues from itzi_core.data_containers import ( - DrainageNodeData, + DrainageLinkAttributes, DrainageLinkData, DrainageNetworkData, - DrainageLinkAttributes, DrainageNodeAttributes, + DrainageNodeData, ) if TYPE_CHECKING: from datetime import datetime + from pyswmm.swmm5 import PySWMM @@ -83,7 +85,7 @@ def __init__( self.swmm_model.swmm_use_hotstart(hotstart_filename) if hotstart_start_datetime is not None: self.swmm_model.setSimulationDateTime( - pyswmm.toolkitapi.SimulationTime.StartDateTime, hotstart_start_datetime + SimulationTime.StartDateTime, hotstart_start_datetime ) self.swmm_model.swmm_start() # allow ponding @@ -200,9 +202,7 @@ def __init__( self.free_weir_coeff = free_weir_coeff self.submerged_weir_coeff = submerged_weir_coeff self.node_type = self.get_node_type() - self.surface_area = self._model.getSimAnalysisSetting( - pyswmm.toolkitapi.SimulationParameters.MinSurfArea - ) + self.surface_area = self._model.getSimAnalysisSetting(SimulationParameters.MinSurfArea) # weir width is the circumference (node considered circular) self.weir_width = 2 * math.sqrt(self.surface_area * math.pi) # Set default values @@ -229,7 +229,7 @@ def get_full_volume(self): return self.surface_area * self.pyswmm_node.full_depth def get_overflow(self): - return self._model.getNodeResult(self.node_id, pyswmm.toolkitapi.NodeResults.overflow) + return self._model.getNodeResult(self.node_id, NodeResults.overflow) def get_crest_elev(self): """Return the crest elevation of the node.""" diff --git a/src/itzi_core/itzi_error.py b/src/itzi_core/itzi_error.py index fdcfcec..be792b9 100644 --- a/src/itzi_core/itzi_error.py +++ b/src/itzi_core/itzi_error.py @@ -16,34 +16,14 @@ class NullError(RuntimeError): """Raised when null values is detected in simulation""" - pass - class DtError(RuntimeError): """Error related to time-step calculation""" - def __init__(self, msg): - self.msg = msg - - def __str__(self): - return repr(self.msg) - class MassBalanceError(RuntimeError): """Raised when mass balance error exceeds threshold""" - def __init__(self, msg: str): - self.msg = msg - - def __str__(self): - return repr(self.msg) - class HotstartError(RuntimeError): """Raised when hotstart file operations fail.""" - - def __init__(self, msg): - self.msg = msg - - def __str__(self): - return repr(self.msg) diff --git a/src/itzi_core/profiler.py b/src/itzi_core/profiler.py index 7210f6e..96d764a 100644 --- a/src/itzi_core/profiler.py +++ b/src/itzi_core/profiler.py @@ -18,8 +18,9 @@ """ import os -from pathlib import Path +from collections.abc import Iterator from contextlib import contextmanager +from pathlib import Path # Attempt to import pyinstrument try: @@ -32,7 +33,7 @@ @contextmanager -def profile_context(file_path: Path = None): +def profile_context(file_path: Path | None = None) -> Iterator[None]: """ A context manager for profiling code blocks. @@ -49,7 +50,9 @@ def profile_context(file_path: Path = None): # Code to be profiled (or not) run_simulation() """ - profiler_active = os.environ.get("ITZI_PROFILE") == "1" and PYINSTRUMENT_AVAILABLE + profiler_active = ( + os.environ.get("ITZI_PROFILE") == "1" and PYINSTRUMENT_AVAILABLE and Profiler is not None + ) if profiler_active: profiler = Profiler() diff --git a/src/itzi_core/providers/csv_output.py b/src/itzi_core/providers/csv_output.py index f11ff05..4e3e38f 100644 --- a/src/itzi_core/providers/csv_output.py +++ b/src/itzi_core/providers/csv_output.py @@ -18,7 +18,7 @@ from datetime import datetime, timedelta from io import StringIO from pathlib import PurePosixPath, PureWindowsPath -from typing import TYPE_CHECKING, TypedDict +from typing import TYPE_CHECKING, Any, TypedDict import pandas as pd @@ -61,10 +61,8 @@ class CSVVectorOutputProvider(VectorOutputProvider): def __init__(self, config: CSVVectorOutputConfig) -> None: """Initialize output provider with provider configuration.""" - try: - self.srid = config["crs"].to_epsg() - except AttributeError: - self.srid = 0 + crs = config["crs"] + self.srid = 0 if crs is None else crs.to_epsg() or 0 self.store = config["store"] prefix_str = config["results_prefix"] prefix_path = PurePosixPath(prefix_str.replace("\\", "/")) @@ -76,11 +74,14 @@ def __init__(self, config: CSVVectorOutputConfig) -> None: raise ValueError("results_prefix must be a relative path without parent traversal") results_prefix = prefix_path.as_posix() if prefix_path.parts else "" - self.existing_ids = {"link": None, "node": None} # Objects ids already in the file - self.existing_max_time = {"link": None, "node": None} # Max of sim_time in existing_file + self.existing_ids: dict[str, set[Any] | None] = {"link": None, "node": None} + self.existing_max_time: dict[str, datetime | timedelta | None] = { + "link": None, + "node": None, + } self.number_of_writes = {"link": 0, "node": 0} - self.file_paths = {"link": None, "node": None} - self.headers = {"link": None, "node": None} + self.file_paths: dict[str, str] = {} + self.headers: dict[str, list[str]] = {} self.append_mode = {"link": True, "node": True} if config["overwrite"]: self.append_mode = {"link": False, "node": False} @@ -97,8 +98,6 @@ def __init__(self, config: CSVVectorOutputConfig) -> None: # create the CSV files if not self.append_mode[geom_type]: self._write_headers(geom_type) - print(self.existing_ids) - print(self.existing_max_time) def write_vector( self, drainage_data: DrainageNetworkData, sim_time: datetime | timedelta @@ -130,7 +129,6 @@ def _check_existing_csv(self, geom_type: str): - new object ID ≠ existing ones could not be checked without drainage network data """ - existing_csv = None try: existing_csv = StringIO( bytes(obstore.get(self.store, self.file_paths[geom_type]).bytes()).decode("utf-8") @@ -143,7 +141,6 @@ def _check_existing_csv(self, geom_type: str): expected_headers = self.headers[geom_type] if not existing_headers == expected_headers: raise ValueError(f"Headers mismatch in existing file {self.file_paths[geom_type]}.") - self.append_mode[geom_type] = False id_col = f"{geom_type}_id" # Store values existing ids @@ -162,7 +159,6 @@ def _check_existing_csv(self, geom_type: str): raise ValueError( f"Unknown sim_time column in existing file {self.file_paths[geom_type]}." ) - print(df_csv) def _write_headers(self, geom_type: str): """Create an in-memory CSV file with headers and save it in the store.""" @@ -178,29 +174,33 @@ def _validate_time_on_first_write(self, sim_time: datetime | timedelta) -> None: for geom_type in ["node", "link"]: # Only validate on first write - if self.number_of_writes[geom_type] > 0 or self.existing_max_time[geom_type] is None: + existing_time = self.existing_max_time[geom_type] + if self.number_of_writes[geom_type] > 0 or existing_time is None: continue - # Type must match - if type(self.existing_max_time[geom_type]) is not type(sim_time): + # ty finds an error when only one if is used + # ruff: noqa: SIM114 + if isinstance(sim_time, datetime) and isinstance(existing_time, datetime): + time_is_increasing = sim_time > existing_time + elif isinstance(sim_time, timedelta) and isinstance(existing_time, timedelta): + time_is_increasing = sim_time > existing_time + else: time_type_name = ( "relative (timedelta)" if isinstance(sim_time, timedelta) else "absolute (datetime)" ) existing_type_name = ( - "relative" - if isinstance(self.existing_max_time[geom_type], timedelta) - else "absolute" + "relative" if isinstance(existing_time, timedelta) else "absolute" ) - raise ValueError( + raise TypeError( f"Time type mismatch for {geom_type}: " f"attempting to write {time_type_name} but existing file has {existing_type_name}" ) # Time must increase - if not sim_time > self.existing_max_time[geom_type]: + if not time_is_increasing: raise ValueError( f"Time not increasing for {geom_type}: attempting to write {sim_time} but " - f"existing file has a max sim_time value of {self.existing_max_time[geom_type]}" + f"existing file has a max sim_time value of {existing_time}" ) def _update_csv( @@ -235,7 +235,7 @@ def _update_csv( updated_csv = existing_csv + new_rows obstore.put(self.store, self.file_paths[geom_type], file=updated_csv.encode("utf-8")) - def _attrs_line(self, drainage_element: DrainageNodeData | DrainageLinkData) -> list[str, ...]: + def _attrs_line(self, drainage_element: DrainageNodeData | DrainageLinkData) -> list[str]: """Return a list of attributes""" # Convert attributes to list attributes = [str(a) for a in drainage_element.attributes.model_dump().values()] diff --git a/src/itzi_core/providers/memory_output.py b/src/itzi_core/providers/memory_output.py index b681dce..538462b 100644 --- a/src/itzi_core/providers/memory_output.py +++ b/src/itzi_core/providers/memory_output.py @@ -29,7 +29,9 @@ def __init__(self, out_map_names: Mapping[str, str]) -> None: """Initialize output provider with simulation configuration.""" # user-selected map names. self.out_map_names = out_map_names - self.output_maps_dict: dict[str, list] = {k: [] for k in self.out_map_names} + self.output_maps_dict: dict[str, list[tuple[datetime | timedelta, np.ndarray]]] = { + key: [] for key in self.out_map_names + } def write_arrays( self, array_dict: Mapping[str, np.ndarray], sim_time: datetime | timedelta diff --git a/src/itzi_core/providers/xarray_input.py b/src/itzi_core/providers/xarray_input.py index d3fd3b0..62852ee 100644 --- a/src/itzi_core/providers/xarray_input.py +++ b/src/itzi_core/providers/xarray_input.py @@ -14,13 +14,14 @@ from __future__ import annotations -from typing import Iterable, Mapping, TypedDict, NotRequired, TYPE_CHECKING +from collections.abc import Iterable, Mapping +from typing import TYPE_CHECKING, NotRequired, TypedDict import numpy as np try: - import xarray as xr import pandas as pd + import xarray as xr except ImportError: raise ImportError( "To use the xarray input backend, install itzi with: " @@ -28,9 +29,9 @@ "or 'pip install itzi[cloud]'" ) +from itzi_core.const import TemporalType from itzi_core.providers.base import RasterInputProvider from itzi_core.providers.domain_data import DomainData -from itzi_core.const import TemporalType if TYPE_CHECKING: from datetime import datetime @@ -148,9 +149,10 @@ def _validate_dimensions(self) -> None: """Validate that: - the specified spatial dimensions exist in the dataset. - The dimensions are one-dimensional.""" - for var_name in self.dataset.data_vars.keys(): + for raw_var_name in self.dataset.data_vars.keys(): + var_name = str(raw_var_name) da_var: xr.DataArray = self.dataset[var_name] - var_dims: set[str] = set(da_var.dims) + var_dims = {str(dim) for dim in da_var.dims} for dim_type in ["x", "y"]: dim_name: str = self.dataset_dims[var_name][dim_type] if dim_name not in var_dims: @@ -169,7 +171,7 @@ def _validate_variables_dimensionality(self) -> None: """Validate that all variables are either 2D[y, x] or 3D[time, y, x].""" for var_name in self.input_map_names.values(): da_var: xr.DataArray = self.dataset[var_name] - var_dims: set[str] = set(da_var.dims) + var_dims = {str(dim) for dim in da_var.dims} num_dims = len(da_var.dims) x_dim: str = self.dataset_dims[var_name]["x"] y_dim: str = self.dataset_dims[var_name]["y"] @@ -213,7 +215,7 @@ def _validate_equal_spacing_of_spatial_dims(self) -> None: diffs = np.diff(coord.values if hasattr(coord, "values") else coord) # no coordinates present if len(diffs) == 0: - pass + continue if not np.allclose(diffs, diffs[0]): raise ValueError( f"Dimension {dim_name} of variable {var_name} not equally spaced." @@ -231,7 +233,7 @@ def _validate_equality_of_spatial_dims(self): if len(dim_names) == 0: continue da_list: list[xr.DataArray] = [self.dataset[dim_name] for dim_name in dim_names] - ref_da: np.ndarray = da_list[0] + ref_da: xr.DataArray = da_list[0] for da in da_list: if not np.allclose(ref_da.values, da.values): raise ValueError( diff --git a/src/itzi_core/rasterdomain.py b/src/itzi_core/rasterdomain.py index 5d58db1..85e1dd3 100644 --- a/src/itzi_core/rasterdomain.py +++ b/src/itzi_core/rasterdomain.py @@ -12,82 +12,15 @@ GNU Lesser General Public License for more details. """ -from datetime import datetime -from typing import Self, Callable, TYPE_CHECKING import io +from typing import Self import numpy as np from itzi_core.array_definitions import ARRAY_DEFINITIONS, ArrayCategory -from .compute import rastermetrics from itzi_core.itzi_error import HotstartError -if TYPE_CHECKING: - from itzi_core.providers.base import RasterInputProvider - - -class TimedArray: - """A container for np.ndarray with time information. - Update the array value according to the simulation time. - array is accessed via get() - """ - - def __init__( - self, - mkey: str, - raster_provider: "RasterInputProvider", - default_array_func: Callable[[], np.ndarray], - ) -> None: - assert isinstance(mkey, str), "not a string!" - assert hasattr(default_array_func, "__call__"), "not a function!" - self.mkey = mkey # An array identifier - self.raster_provider = raster_provider - # A function to generate a default array - self.default_array_func = default_array_func - # default values for start and end - # intended to trigger update when is_valid() is first called - self.arr_start = datetime(1, 1, 2) - self.arr_end = datetime(1, 1, 1) - # Necessary for BMI implementation - self.origin = raster_provider.get_origin() - # Placeholder for the numpy array - self.arr = None - - def get(self, sim_time: datetime) -> np.ndarray: - """Return a numpy array valid for the given time - If the array stored is not valid, update the values of the object - """ - assert isinstance(sim_time, datetime), "not a datetime object!" - if not self.is_valid(sim_time): - self.update_values(sim_time) - return self.arr - - def is_valid(self, sim_time: datetime) -> bool: - """input being a time in datetime - If the current stored array is within the half-open range [start, end), - return True - If not return False - """ - return bool(self.arr_start <= sim_time < self.arr_end) - - def update_values(self, sim_time: datetime) -> Self: - """Update array, start_time and end_time from provider - if the provider returns None, set array to default value - """ - # Retrieve values - arr, arr_start, arr_end = self.raster_provider.get_array(self.mkey, sim_time) - # set to default if no array retrieved - if not isinstance(arr, np.ndarray): - arr = self.default_array_func() - # check retrieved values - assert isinstance(arr_start, datetime), "not a datetime object!" - assert isinstance(arr_end, datetime), "not a datetime object!" - assert arr_start <= sim_time < arr_end, "wrong time retrieved!" - # update object values - self.arr_start = arr_start - self.arr_end = arr_end - self.arr = arr - return self +from .compute import rastermetrics class RasterDomain: @@ -132,8 +65,8 @@ def __init__(self, dtype, arr_mask: np.ndarray, cell_shape: tuple[float, float]) if arr_def.key in self.k_all } # Instantiate arrays and padded arrays filled with zeros - self.arr = dict.fromkeys(self.k_all) - self.arrp = dict.fromkeys(self.k_all) + self.arr: dict[str, np.ndarray] = {} + self.arrp: dict[str, np.ndarray] = {} self._create_arrays() def pad_array(self, arr) -> tuple[np.ndarray, np.ndarray]: @@ -148,14 +81,13 @@ def _create_arrays(self) -> Self: """Instantiate masked arrays and padded arrays the unpadded arrays are a slice of the padded ones """ - for k in self.arr.keys(): + for k in self.k_all: arr = np.full(fill_value=self.fill_values[k], shape=self.shape, dtype=self.dtypes[k]) self.arr[k], self.arrp[k] = self.pad_array(arr) return self def update_mask(self, arr: np.ndarray) -> Self: """Create a mask array by marking NULL values from arr as True.""" - pass # self.mask[:] = np.isnan(arr) return self diff --git a/src/itzi_core/report.py b/src/itzi_core/report.py index 60980ec..f78ed81 100644 --- a/src/itzi_core/report.py +++ b/src/itzi_core/report.py @@ -200,7 +200,7 @@ def write_mass_balance(self, data: SimulationData, converted_sim_time: datetime closure_residual, relative_closure_error = calculate_closure( continuity_data.volume_change, signed_volume_terms, - active_cells * cell_area, + float(active_cells * cell_area), ) # 3. Assemble data and log diff --git a/src/itzi_core/simulation.py b/src/itzi_core/simulation.py index 2fe445a..a5286c9 100644 --- a/src/itzi_core/simulation.py +++ b/src/itzi_core/simulation.py @@ -42,9 +42,10 @@ from itzi_core.drainage import DrainageSimulation from itzi_core.hydrology import Hydrology from itzi_core.providers.domain_data import DomainData - from itzi_core.rasterdomain import RasterDomain, TimedArray + from itzi_core.rasterdomain import RasterDomain from itzi_core.report import Report from itzi_core.surfaceflow import SurfaceFlowSimulation + from itzi_core.timed_array import TimedArraySource logger = logging.getLogger(__name__) @@ -64,7 +65,7 @@ def __init__( hydrology_model: Hydrology, surface_flow: SurfaceFlowSimulation, drainage_model: DrainageSimulation | None, - drainage_nodes_list: list[DrainageNodeCouplingData] | None, + drainage_nodes_list: list[DrainageNodeCouplingData], report: Report, ): self.sim_config = sim_config @@ -92,12 +93,13 @@ def __init__( for arr_def in ARRAY_DEFINITIONS if arr_def.computes_from is not None and ArrayCategory.ACCUMULATION in arr_def.category } - self.accum_update_time: dict[str, datetime | None] = { - accum: None for source, accum in self.accum_mapping.items() + self.accum_update_time: dict[str, datetime] = { + accum: self.sim_time for accum in self.accum_mapping.values() } + self._initialized = False + self.node_id_to_loc: dict[str, tuple[int, int]] = {} if self.drainage_model: - assert self.drainage_nodes_list is not None - self.node_id_to_loc: dict[str, tuple[int, int]] = { + self.node_id_to_loc = { n.node_id: (n.row, n.col) for n in self.drainage_nodes_list if n.node_object.is_coupled() and n.row is not None and n.col is not None @@ -141,7 +143,7 @@ def end_time(self) -> datetime: return self.schedule.end_time @property - def timed_arrays(self) -> dict[str, TimedArray] | None: + def timed_arrays(self) -> dict[str, TimedArraySource] | None: if self.timed_input_manager is None: return None return self.timed_input_manager.timed_arrays @@ -167,6 +169,7 @@ def initialize(self) -> Self: self.raster_domain.reset_accumulations() for key in self.accum_update_time: self.accum_update_time[key] = self.sim_time + self._initialized = True return self def update(self) -> Self: @@ -422,9 +425,6 @@ def _update_accum_array(self, k: str, sim_time: datetime) -> None: """ ak = self.accum_mapping[k] last_update = self.accum_update_time[ak] - if last_update is None: - self.accum_update_time[ak] = sim_time - return time_diff = (sim_time - last_update).total_seconds() if time_diff > 0: rate_array = self.raster_domain.get_padded(k) @@ -444,17 +444,11 @@ def create_hotstart(self) -> io.BytesIO: Raises: RuntimeError: If called before initialize() has established valid state. """ - # Guard: Check that initialize() has been called - # accum_update_time values are set to None in __init__ and only - # become valid datetime objects after initialize() is called - if any(v is None for v in self.accum_update_time.values()): + if not self._initialized: raise RuntimeError( "Cannot create hotstart: simulation has not been initialized. " "Call initialize() before creating a hotstart." ) - accum_update_time = { - key: value for key, value in self.accum_update_time.items() if value is not None - } # Get SWMM hotstart bytes if drainage is enabled swmm_hotstart_bytes: bytes | None = None @@ -473,7 +467,7 @@ def create_hotstart(self) -> io.BytesIO: dt=self.dt.total_seconds(), next_ts=self.schedule.snapshot_deadlines(), time_steps_counters=self.time_steps_counters, - accum_update_time=accum_update_time, + accum_update_time=dict(self.accum_update_time), old_domain_volume=self.old_domain_volume, swmm_elapsed_time=swmm_elapsed_time, ) @@ -550,6 +544,7 @@ def restore_state(self, simulation_state: HotstartSimulationState) -> Self: # Restore accumulation update timestamps self.accum_update_time = dict(simulation_state.accum_update_time) + self._initialized = True # Restore old domain volume for continuity tracking self.old_domain_volume = simulation_state.old_domain_volume diff --git a/src/itzi_core/simulation_builder.py b/src/itzi_core/simulation_builder.py index 3744b31..3bbce2c 100644 --- a/src/itzi_core/simulation_builder.py +++ b/src/itzi_core/simulation_builder.py @@ -16,15 +16,16 @@ import io import tempfile +from collections.abc import Iterable from datetime import timedelta from pathlib import Path -from typing import TYPE_CHECKING, ClassVar +from typing import TYPE_CHECKING, Any, ClassVar import numpy as np import pyswmm from numpy.typing import ArrayLike, DTypeLike -from itzi_core import infiltration, rasterdomain +from itzi_core import infiltration from itzi_core.array_definitions import ARRAY_DEFINITIONS, ArrayCategory from itzi_core.const import InfiltrationModelType from itzi_core.data_containers import DrainageNodeCouplingData @@ -32,11 +33,13 @@ from itzi_core.hotstart import HotstartLoader from itzi_core.hydrology import Hydrology from itzi_core.itzi_error import HotstartError +from itzi_core.rasterdomain import RasterDomain from itzi_core.report import Report from itzi_core.simulation import Simulation from itzi_core.simulation_schedule import SimulationSchedule from itzi_core.surfaceflow import SurfaceFlowSimulation from itzi_core.swmm_input_parser import SwmmInputParser +from itzi_core.timed_array import TimedArray from itzi_core.timed_inputs import TimedInputManager if TYPE_CHECKING: @@ -83,9 +86,9 @@ def __init__( sim_config: SimulationConfig, arr_mask: ArrayLike, dtype: DTypeLike = np.float32, - ): + ) -> None: self.sim_config = sim_config - self.arr_mask = arr_mask + self.arr_mask = np.asarray(arr_mask) self.dtype = dtype # Optional components (set via builder methods) @@ -152,7 +155,7 @@ def with_mass_balance_output_provider( self.mass_balance_output_provider = provider return self - def _validate_hotstart_congruence(self) -> None: + def _validate_hotstart_congruence(self, hotstart_loader: HotstartLoader) -> None: """Validate hotstart data against builder configuration. This method performs congruence checks between the hotstart metadata @@ -163,9 +166,9 @@ def _validate_hotstart_congruence(self) -> None: HotstartError: If any congruence check fails. """ - hotstart_domain = self.hotstart_loader.get_domain_data() - hotstart_config = self.hotstart_loader.get_simulation_config() - hotstart_state = self.hotstart_loader.get_simulation_state() + hotstart_domain = hotstart_loader.get_domain_data() + hotstart_config = hotstart_loader.get_simulation_config() + hotstart_state = hotstart_loader.get_simulation_state() # Validate domain metadata self._validate_domain_congruence(hotstart_domain) @@ -174,7 +177,7 @@ def _validate_hotstart_congruence(self) -> None: self._validate_mask_congruence(hotstart_domain) # Validate drainage expectations - self._validate_drainage_congruence(hotstart_config) + self._validate_drainage_congruence(hotstart_config, hotstart_loader) # Validate resume-time configuration compatibility self._validate_resume_config_congruence(hotstart_config, hotstart_state) @@ -322,7 +325,11 @@ def _validate_mask_congruence(self, hotstart_domain: DomainData) -> None: f"hotstart expects {expected_shape}" ) - def _validate_drainage_congruence(self, hotstart_config: SimulationConfig) -> None: + def _validate_drainage_congruence( + self, + hotstart_config: SimulationConfig, + hotstart_loader: HotstartLoader, + ) -> None: """Validate drainage expectations match between hotstart and current config.""" hotstart_has_drainage = hotstart_config.swmm_inp is not None builder_has_drainage = self.sim_config.swmm_inp is not None @@ -338,30 +345,35 @@ def _validate_drainage_congruence(self, hotstart_config: SimulationConfig) -> No ) # If both have drainage, check that SWMM hotstart bytes are present - if hotstart_has_drainage and not self.hotstart_loader.has_swmm_hotstart(): + if hotstart_has_drainage and not hotstart_loader.has_swmm_hotstart(): raise HotstartError( "Hotstart metadata indicates drainage but SWMM hotstart file is missing from archive" ) def build(self) -> Simulation: """Build a simulation and explicitly load or prime provider-backed inputs.""" - # Validate required components - if self.domain_data is None: + domain_data = self.domain_data + raster_output_provider = self.raster_output_provider + vector_output_provider = self.vector_output_provider + input_provider = self.raster_input_provider + hotstart_loader = self.hotstart_loader + + if domain_data is None: raise ValueError("Domain data must be set via input provider or directly") - if self.raster_output_provider is None or self.vector_output_provider is None: + if raster_output_provider is None or vector_output_provider is None: raise ValueError("Output providers are mandatory") # Validate hotstart congruence before building - if self.hotstart_loader: - self._validate_hotstart_congruence() + if hotstart_loader is not None: + self._validate_hotstart_congruence(hotstart_loader) # Create timed arrays if input provider exists timed_arrays = None - if self.raster_input_provider: - timed_arrays = self._create_timed_arrays() + if input_provider is not None: + timed_arrays = self._create_timed_arrays(input_provider, domain_data) # Create raster domain - raster_domain = self._create_raster_domain(self.domain_data.cell_shape) + raster_domain = self._create_raster_domain(domain_data.cell_shape) # Create models infiltration_model = self._create_infiltration_model(raster_domain) @@ -371,7 +383,7 @@ def build(self) -> Simulation: ) # Create drainage with optional SWMM hotstart injection - nodes_list, drainage_sim = self._create_drainage_simulation() + nodes_list, drainage_sim = self._create_drainage_simulation(domain_data) schedule = SimulationSchedule( self.sim_config.start_time, self.sim_config.end_time, @@ -391,8 +403,8 @@ def build(self) -> Simulation: report = Report( start_time=self.sim_config.start_time, temporal_type=self.sim_config.temporal_type, - raster_output_provider=self.raster_output_provider, - vector_output_provider=self.vector_output_provider, + raster_output_provider=raster_output_provider, + vector_output_provider=vector_output_provider, mass_balance_output_provider=self.mass_balance_output_provider, out_map_names=self.sim_config.output_map_names, dt=self.sim_config.record_step, @@ -401,7 +413,7 @@ def build(self) -> Simulation: # Create simulation simulation = Simulation( self.sim_config, - self.domain_data, + domain_data, raster_domain, schedule, timed_input_manager, @@ -413,13 +425,13 @@ def build(self) -> Simulation: ) # Apply hotstart restore if hotstart data is present - if self.hotstart_loader: - raster_state_buffer = self.hotstart_loader.get_raster_state_buffer() + if hotstart_loader is not None: + raster_state_buffer = hotstart_loader.get_raster_state_buffer() raster_domain.load_state(raster_state_buffer) - simulation_state = self.hotstart_loader.get_simulation_state() + simulation_state = hotstart_loader.get_simulation_state() simulation.restore_state(simulation_state) - hotstart_config = self.hotstart_loader.get_simulation_config() + hotstart_config = hotstart_loader.get_simulation_config() changed_input_keys = self._changed_input_keys(hotstart_config) restored_input_deadline = simulation.schedule.deadline("input") restored_end_deadline = simulation.schedule.deadline("end") @@ -462,27 +474,29 @@ def build(self) -> Simulation: return simulation - def _create_timed_arrays(self) -> dict[str, rasterdomain.TimedArray]: + def _create_timed_arrays( + self, + input_provider: RasterInputProvider, + domain_data: DomainData, + ) -> dict[str, TimedArray]: """Create configured time-varying raster inputs.""" timed_arrays = {} input_keys = [ arr_def.key for arr_def in ARRAY_DEFINITIONS if ArrayCategory.INPUT in arr_def.category ] - raster_shape = (self.domain_data.rows, self.domain_data.cols) + raster_shape = (domain_data.rows, domain_data.cols) def zeros_array_func() -> np.ndarray: return np.zeros(shape=raster_shape, dtype=self.dtype) for arr_key in input_keys: - timed_arrays[arr_key] = rasterdomain.TimedArray( - arr_key, self.raster_input_provider, zeros_array_func - ) + timed_arrays[arr_key] = TimedArray(arr_key, input_provider, zeros_array_func) return timed_arrays - def _create_raster_domain(self, cell_shape) -> rasterdomain.RasterDomain: + def _create_raster_domain(self, cell_shape) -> RasterDomain: """Create a raster domain.""" try: - raster_domain = rasterdomain.RasterDomain( + raster_domain = RasterDomain( dtype=self.dtype, arr_mask=self.arr_mask, cell_shape=cell_shape, @@ -493,7 +507,7 @@ def _create_raster_domain(self, cell_shape) -> rasterdomain.RasterDomain: def _create_infiltration_model( self, - raster_domain: rasterdomain.RasterDomain, + raster_domain: RasterDomain, ) -> infiltration.InfiltrationModel: """Create an infiltration model based on configuration.""" inf_model = self.sim_config.infiltration_model @@ -512,7 +526,8 @@ def _create_infiltration_model( def _create_drainage_simulation( self, - ) -> tuple[list[DrainageNodeCouplingData] | None, DrainageSimulation | None]: + domain_data: DomainData, + ) -> tuple[list[DrainageNodeCouplingData], DrainageSimulation | None]: """Create drainage simulation components if SWMM input is provided. If hotstart data includes SWMM state, writes the SWMM hotstart bytes @@ -525,7 +540,7 @@ def _create_drainage_simulation( in their timeseries rather than restarting from T=0. """ if not self.sim_config.swmm_inp: - return None, None + return [], None swmm_input_path = str(self.sim_config.swmm_inp) @@ -550,6 +565,7 @@ def _create_drainage_simulation( nodes_list: list[DrainageNodeCouplingData] = self._get_nodes_list( all_nodes, nodes_coors_dict, + domain_data=domain_data, orifice_coeff=self.sim_config.orifice_coeff, free_weir_coeff=self.sim_config.free_weir_coeff, submerged_weir_coeff=self.sim_config.submerged_weir_coeff, @@ -562,8 +578,10 @@ def _create_drainage_simulation( node_objects_only = [i.node_object for i in nodes_list] # Handle SWMM hotstart injection if present - if self.hotstart_loader and self.hotstart_loader.has_swmm_hotstart(): + if self.hotstart_loader is not None and self.hotstart_loader.has_swmm_hotstart(): swmm_bytes = self.hotstart_loader.get_swmm_hotstart_bytes() + if swmm_bytes is None: + raise HotstartError("SWMM hotstart data is missing from the archive") # Create a temporary file for SWMM to read. # delete_on_close=False keeps the file after closing so SWMM can open it. with tempfile.NamedTemporaryFile(suffix=".hsf", delete_on_close=False) as tmp: @@ -584,15 +602,16 @@ def _create_drainage_simulation( def _get_nodes_list( self, - pswmm_nodes: list, - nodes_coor_dict: dict, + pswmm_nodes: Iterable[Any], + nodes_coor_dict: dict[str, Any], + domain_data: DomainData, orifice_coeff: float, free_weir_coeff: float, submerged_weir_coeff: float, g: float, ) -> list[DrainageNodeCouplingData]: """Check if the drainage nodes are inside the region and can be coupled. - Return a list of DrainageNodeCouplingData + A node without coordinates cannot be coupled. """ nodes_list = [] for pyswmm_node in pswmm_nodes: @@ -606,8 +625,10 @@ def _get_nodes_list( submerged_weir_coeff=submerged_weir_coeff, g=g, ) - # a node without coordinates cannot be coupled - if coors is None or not self.domain_data.is_in_domain(x=coors.x, y=coors.y): + pixel = ( + None if coors is None else domain_data.coordinates_to_pixel(x=coors.x, y=coors.y) + ) + if pixel is None: x_coor = None y_coor = None row = None @@ -617,7 +638,7 @@ def _get_nodes_list( node.coupling_type = CouplingTypes.COUPLED_NO_FLOW x_coor = coors.x y_coor = coors.y - row, col = self.domain_data.coordinates_to_pixel(x=x_coor, y=y_coor) + row, col = pixel # populate list drainage_node_data = DrainageNodeCouplingData( node_id=pyswmm_node.nodeid, node_object=node, x=x_coor, y=y_coor, row=row, col=col diff --git a/src/itzi_core/surfaceflow.py b/src/itzi_core/surfaceflow.py index 8663060..9f6fd85 100644 --- a/src/itzi_core/surfaceflow.py +++ b/src/itzi_core/surfaceflow.py @@ -109,9 +109,7 @@ def dt(self): def dt(self, newdt: timedelta): """return an error if new dt is higher than current one or negative""" newdt_s = newdt.total_seconds() - if self._dt is None: - self._dt = newdt_s - elif newdt_s <= 0: + if newdt_s <= 0: raise DtError(f"dt must be positive, not {newdt_s}s") elif newdt_s > self._dt + self._dt_fudge: raise DtError( diff --git a/src/itzi_core/swmm_input_parser.py b/src/itzi_core/swmm_input_parser.py index afa2321..3c88954 100644 --- a/src/itzi_core/swmm_input_parser.py +++ b/src/itzi_core/swmm_input_parser.py @@ -15,13 +15,14 @@ import os from collections import namedtuple from datetime import datetime +from typing import ClassVar -class SwmmInputParser(object): +class SwmmInputParser: """A parser for swmm input text file""" # list of sections keywords - sections_kwd = [ + sections_kwd: ClassVar[tuple[str, ...]] = ( "title", # project title "option", # analysis options "junction", # junction node information @@ -36,10 +37,10 @@ class SwmmInputParser(object): "xsection", # conduit, orifice, and weir cross-section geometry "coordinate", # coordinates of drainage system nodes "vertice", # coordinates of interior vertex points of links - ] - link_types = ["conduit", "pump", "orifice", "weir", "outlet"] + ) + link_types: ClassVar[tuple[str, ...]] = ("conduit", "pump", "orifice", "weir", "outlet") # define object containers - junction_values = ["x", "y", "elev", "ymax", "y0", "ysur", "apond"] + junction_values: ClassVar[tuple[str, ...]] = ("x", "y", "elev", "ymax", "y0", "ysur", "apond") Junction = namedtuple("Junction", junction_values) Link = namedtuple("Link", ["in_node", "out_node", "vertices"]) # coordinates container @@ -48,7 +49,7 @@ class SwmmInputParser(object): def __init__(self, input_file): # read and parse the input file assert os.path.isfile(input_file) - self.inp = dict.fromkeys(self.sections_kwd) + self.inp: dict[str, list[list[str]]] = {section: [] for section in self.sections_kwd} self.read_inp(input_file) def section_kwd(self, sect_name): @@ -77,8 +78,6 @@ def read_inp(self, input_file): elif current_section is None: continue else: - if self.inp[current_section] is None: - self.inp[current_section] = [] self.inp[current_section].append(line.strip().split()) def get_juntions_ids(self): @@ -130,22 +129,20 @@ def get_links_id_as_dict(self): # loop through all types of links for k in self.link_types: links = self.inp[k] - if links is not None: - for ln in links: - ID = ln[0] - vertices = self.get_vertices(ID) - # names of link, inlet and outlet nodes - links_dict[ID] = self.Link(in_node=ln[1], out_node=ln[2], vertices=vertices) + for ln in links: + ID = ln[0] + vertices = self.get_vertices(ID) + # names of link, inlet and outlet nodes + links_dict[ID] = self.Link(in_node=ln[1], out_node=ln[2], vertices=vertices) return links_dict def get_vertices(self, link_name): """For a given link name, return a list of Coordinates objects""" vertices = [] - if isinstance(self.inp["vertice"], list): - for vertex in self.inp["vertice"]: - if link_name == vertex[0]: - vertex_c = self.Coordinates(float(vertex[1]), float(vertex[2])) - vertices.append(vertex_c) + for vertex in self.inp["vertice"]: + if link_name == vertex[0]: + vertex_c = self.Coordinates(float(vertex[1]), float(vertex[2])) + vertices.append(vertex_c) return vertices def get_start_datetime(self) -> datetime | None: diff --git a/src/itzi_core/timed_array.py b/src/itzi_core/timed_array.py new file mode 100644 index 0000000..cac20a6 --- /dev/null +++ b/src/itzi_core/timed_array.py @@ -0,0 +1,90 @@ +""" +Copyright (C) 2026 Laurent G. Courty + +This library is free software; you can redistribute it and/or +modify it under the terms of the GNU Lesser General Public License +as published by the Free Software Foundation; either version 2.1 +of the License, or (at your option) any later version. + +This library is distributed in the hope that it will be useful, +but WITHOUT ANY WARRANTY; without even the implied warranty of +MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the +GNU Lesser General Public License for more details. +""" + +from __future__ import annotations + +from collections.abc import Callable +from datetime import datetime +from typing import TYPE_CHECKING, Protocol, Self + +import numpy as np + +if TYPE_CHECKING: + from itzi_core.providers.base import RasterInputProvider + + +class TimedArraySource(Protocol): + arr_start: datetime + arr_end: datetime + arr: np.ndarray | None + + def is_valid(self, sim_time: datetime) -> bool: ... + + def get(self, sim_time: datetime) -> np.ndarray: ... + + +class TimedArray: + """A container for np.ndarray with time information. + Update the array value according to the simulation time. + array is accessed via get() + """ + + def __init__( + self, + mkey: str, + raster_provider: RasterInputProvider, + default_array_func: Callable[[], np.ndarray], + ) -> None: + self.mkey = mkey # An array identifier + self.raster_provider = raster_provider + # A function to generate a default array + self.default_array_func = default_array_func + # default values for start and end + # intended to trigger update when is_valid() is first called + self.arr_start = datetime(1, 1, 2) + self.arr_end = datetime(1, 1, 1) + # Necessary for BMI implementation + self.origin = raster_provider.get_origin() + self.arr: np.ndarray = self.default_array_func() + + def get(self, sim_time: datetime) -> np.ndarray: + """Return a numpy array valid for the given time + If the array stored is not valid, update the values of the object + """ + if not self.is_valid(sim_time): + self._update_values(sim_time) + return self.arr + + def is_valid(self, sim_time: datetime) -> bool: + """input being a time in datetime + If the current stored array is within the half-open range [start, end), + return True + If not return False + """ + return bool(self.arr_start <= sim_time < self.arr_end) + + def _update_values(self, sim_time: datetime) -> Self: + """Update array, start_time and end_time from provider + if the provider returns None, set array to default value + """ + # Retrieve values + arr, arr_start, arr_end = self.raster_provider.get_array(self.mkey, sim_time) + # set to default if no array retrieved + if arr is None: + arr = self.default_array_func() + # update object values + self.arr_start = arr_start + self.arr_end = arr_end + self.arr = arr + return self diff --git a/src/itzi_core/timed_inputs.py b/src/itzi_core/timed_inputs.py index 90fc016..7efcece 100644 --- a/src/itzi_core/timed_inputs.py +++ b/src/itzi_core/timed_inputs.py @@ -15,12 +15,16 @@ from __future__ import annotations import logging +from collections.abc import Mapping from datetime import datetime +from typing import TYPE_CHECKING import numpy as np from itzi_core.itzi_error import NullError -from itzi_core.rasterdomain import TimedArray + +if TYPE_CHECKING: + from itzi_core.timed_array import TimedArraySource logger = logging.getLogger(__name__) @@ -33,12 +37,12 @@ class TimedInputManager: def __init__( self, - timed_arrays: dict[str, TimedArray], + timed_arrays: Mapping[str, TimedArraySource], input_wse: bool, end_time: datetime, mask: np.ndarray, ) -> None: - self.timed_arrays = timed_arrays + self.timed_arrays = dict(timed_arrays) self.input_wse = input_wse self.end_time = end_time self.mask = mask diff --git a/tests/ea8b/conftest.py b/tests/ea8b/conftest.py index ff4e110..3b990b0 100644 --- a/tests/ea8b/conftest.py +++ b/tests/ea8b/conftest.py @@ -37,7 +37,7 @@ from itzi_core.const import TemporalType from itzi_core.data_containers import SimulationConfig, SurfaceFlowParameters from itzi_core.providers.csv_mass_balance_output import CSVMassBalanceOutputProvider -from itzi_core.providers.csv_output import CSVVectorOutputProvider +from itzi_core.providers.csv_output import CSVVectorOutputConfig, CSVVectorOutputProvider from itzi_core.providers.icechunk_output import IcechunkRasterOutputProvider from itzi_core.providers.xarray_input import XarrayRasterInputProvider from itzi_core.simulation_builder import SimulationBuilder @@ -191,15 +191,16 @@ def ea8b_simulation(ea8b_data, test_data_path, ea8b_temp_path): ) obj_store = obstore.store.MemoryStore() - vector_output_provider = CSVVectorOutputProvider( - { - "crs": crs, - "store": obj_store, - "results_prefix": "", - "drainage_results_name": sim_config.drainage_output, - "overwrite": True, - } - ) + drainage_results_name = sim_config.drainage_output + assert drainage_results_name is not None + vector_config: CSVVectorOutputConfig = { + "crs": crs, + "store": obj_store, + "results_prefix": "", + "drainage_results_name": drainage_results_name, + "overwrite": True, + } + vector_output_provider = CSVVectorOutputProvider(vector_config) simulation = ( SimulationBuilder(sim_config, arr_mask) @@ -233,7 +234,8 @@ def ea8b_simulation(ea8b_data, test_data_path, ea8b_temp_path): final_state = {} for key in simulation.raster_domain.k_all: final_state[f"raster_{key}"] = simulation.raster_domain.get_array(key) - np.savez(final_state_path, **final_state) + # Keys are generated with a raster_ prefix, so none can alias NumPy's allow_pickle keyword. + np.savez(final_state_path, **final_state) # ty: ignore[invalid-argument-type] return { "obj_store": obj_store, diff --git a/tests/ea8b/helpers.py b/tests/ea8b/helpers.py index 56e5921..73bcda0 100644 --- a/tests/ea8b/helpers.py +++ b/tests/ea8b/helpers.py @@ -12,18 +12,17 @@ GNU Lesser General Public License for more details. """ +import icechunk import numpy as np +import obstore import pandas as pd import pyproj -import icechunk -import obstore from itzi_core.data_containers import SimulationConfig -from itzi_core.simulation_builder import SimulationBuilder -from itzi_core.providers.csv_output import CSVVectorOutputProvider +from itzi_core.providers.csv_output import CSVVectorOutputConfig, CSVVectorOutputProvider from itzi_core.providers.icechunk_output import IcechunkRasterOutputProvider from itzi_core.providers.xarray_input import XarrayRasterInputProvider - +from itzi_core.simulation_builder import SimulationBuilder EA8B_REFERENCE_MIN_NSE = 0.99 EA8B_REFERENCE_MAX_RSR = 0.01 @@ -104,15 +103,16 @@ def build_resumed_simulation( ) obj_store = obstore.store.MemoryStore() - vector_output_provider = CSVVectorOutputProvider( - { - "crs": crs, - "store": obj_store, - "results_prefix": "", - "drainage_results_name": sim_config.drainage_output, - "overwrite": True, - } - ) + drainage_results_name = sim_config.drainage_output + assert drainage_results_name is not None + vector_config: CSVVectorOutputConfig = { + "crs": crs, + "store": obj_store, + "results_prefix": "", + "drainage_results_name": drainage_results_name, + "overwrite": True, + } + vector_output_provider = CSVVectorOutputProvider(vector_config) simulation = ( SimulationBuilder(sim_config, arr_mask) diff --git a/tests/test_5by5.py b/tests/test_5by5.py index 358b834..ae52c6d 100644 --- a/tests/test_5by5.py +++ b/tests/test_5by5.py @@ -75,7 +75,13 @@ def _build_diagnostic_simulation( return simulation -def _run_diagnostic_regression(domain_5by5, helpers, *, force_all: bool): +def _run_diagnostic_regression( + domain_5by5, + helpers, + monkeypatch: pytest.MonkeyPatch, + *, + force_all: bool, +): simulation = _build_diagnostic_simulation( domain_5by5, helpers, @@ -88,16 +94,18 @@ def _run_diagnostic_regression(domain_5by5, helpers, *, force_all: bool): if force_all: scheduled_step = simulation.surface_flow.step - def always_compute_step(*, compute_vdir: bool, compute_froude: bool): + def always_compute_step(*, compute_vdir: bool = True, compute_froude: bool = True): return scheduled_step(compute_vdir=True, compute_froude=True) - simulation.surface_flow.step = always_compute_step + monkeypatch.setattr(simulation.surface_flow, "step", always_compute_step) simulation.initialize() while simulation.sim_time < simulation.end_time: simulation.update() - output_maps = simulation.report.raster_provider.output_maps_dict + raster_provider = simulation.report.raster_provider + assert isinstance(raster_provider, MemoryRasterOutputProvider) + output_maps = raster_provider.output_maps_dict result = { "outputs": { key: [(time, array.copy()) for time, array in output_maps[key]] @@ -187,6 +195,7 @@ def _run_center_pulse_simulation( def test_scheduler_computes_only_requested_report_diagnostics( domain_5by5, helpers, + monkeypatch: pytest.MonkeyPatch, diagnostic_keys: list[str], report_flags: tuple[bool, bool], ): @@ -202,7 +211,7 @@ def test_scheduler_computes_only_requested_report_diagnostics( scheduled_step = simulation.surface_flow.step calls = [] - def tracked_step(*, compute_vdir: bool, compute_froude: bool): + def tracked_step(*, compute_vdir: bool = True, compute_froude: bool = True): step_end = simulation.sim_time + simulation.dt vdir_before = simulation.get_array("vdir").copy() froude_before = simulation.get_array("froude").copy() @@ -222,7 +231,7 @@ def tracked_step(*, compute_vdir: bool, compute_froude: bool): ) return result - simulation.surface_flow.step = tracked_step + monkeypatch.setattr(simulation.surface_flow, "step", tracked_step) simulation.initialize() simulation.get_array("vdir").fill(-123.0) simulation.get_array("froude").fill(-456.0) @@ -238,7 +247,9 @@ def tracked_step(*, compute_vdir: bool, compute_froude: bool): (timedelta(seconds=10), report_flags), ] - output_maps = simulation.report.raster_provider.output_maps_dict + raster_provider = simulation.report.raster_provider + assert isinstance(raster_provider, MemoryRasterOutputProvider) + output_maps = raster_provider.output_maps_dict expected_report_times = [ timedelta(seconds=0), timedelta(seconds=4), @@ -248,12 +259,16 @@ def tracked_step(*, compute_vdir: bool, compute_froude: bool): assert [time for time, _ in output_maps["water_depth"]] == expected_report_times for key in ("vdir", "froude"): expected_times = expected_report_times if key in diagnostic_keys else [] - assert [time for time, _ in output_maps[key]] == expected_times + assert [time for time, _ in output_maps.get(key, [])] == expected_times -def test_lazy_diagnostic_reports_match_always_compute_reference(domain_5by5, helpers): - optimized = _run_diagnostic_regression(domain_5by5, helpers, force_all=False) - reference = _run_diagnostic_regression(domain_5by5, helpers, force_all=True) +def test_lazy_diagnostic_reports_match_always_compute_reference( + domain_5by5, + helpers, + monkeypatch: pytest.MonkeyPatch, +): + optimized = _run_diagnostic_regression(domain_5by5, helpers, monkeypatch, force_all=False) + reference = _run_diagnostic_regression(domain_5by5, helpers, monkeypatch, force_all=True) assert optimized["steps"] == reference["steps"] for key in ("water_depth", "v", "vmax"): diff --git a/tests/test_csv_vector_output.py b/tests/test_csv_vector_output.py index b6f96fe..e52c107 100644 --- a/tests/test_csv_vector_output.py +++ b/tests/test_csv_vector_output.py @@ -13,8 +13,8 @@ """ from datetime import datetime, timedelta -from pathlib import Path from io import StringIO +from pathlib import Path import numpy as np import pytest @@ -23,22 +23,24 @@ pytest.importorskip("pyproj") pytest.importorskip("obstore") -import pyproj import obstore import pandas as pd +import pyproj from itzi_core.data_containers import ( - DrainageNodeAttributes, DrainageLinkAttributes, - DrainageNodeData, DrainageLinkData, DrainageNetworkData, + DrainageNodeAttributes, + DrainageNodeData, ) -from itzi_core.providers.csv_output import CSVVectorOutputConfig, CSVVectorOutputProvider from itzi_core.drainage import CouplingTypes - -from tests.fixtures_vector_output import create_dummy_drainage_network -from tests.fixtures_vector_output import expected_node_coords, expected_vertices +from itzi_core.providers.csv_output import CSVVectorOutputConfig, CSVVectorOutputProvider +from tests.fixtures_vector_output import ( + create_dummy_drainage_network, + expected_node_coords, + expected_vertices, +) @pytest.fixture @@ -94,7 +96,7 @@ def test_csv_no_geom_no_srid(self, results_prefix, sim_time): drainage_network = create_dummy_drainage_network(with_coords=False) obj_store = obstore.store.MemoryStore() file_prefix = "test_drainage_no_geom" - provider_config = { + provider_config: CSVVectorOutputConfig = { "crs": None, "store": obj_store, "results_prefix": results_prefix, @@ -306,7 +308,7 @@ def test_append_success(self, results_prefix): sim_time_3 = timedelta(seconds=120) obj_store = obstore.store.MemoryStore() file_prefix = "test_append_success" - provider_config = { + provider_config: CSVVectorOutputConfig = { "crs": pyproj.CRS.from_epsg(6372), "store": obj_store, "results_prefix": results_prefix, @@ -349,7 +351,7 @@ def test_append_column_mismatch_nodes_to_links(self, test_data_temp_path): file_prefix = "test_column_mismatch" results_prefix = "data" # First, write links data - provider_config = { + provider_config: CSVVectorOutputConfig = { "crs": pyproj.CRS.from_epsg(6372), "store": obj_store, "results_prefix": results_prefix, @@ -372,11 +374,11 @@ def test_append_column_mismatch_nodes_to_links(self, test_data_temp_path): csv_provider_append.write_vector(drainage_network, timedelta(seconds=60)) def test_append_time_type_mismatch_timedelta_to_datetime(self, results_prefix): - """Verify ValueError when appending datetime to file with timedelta.""" + """Verify TypeError when appending datetime to file with timedelta.""" drainage_network = create_dummy_drainage_network() # Write with timedelta - provider_config = { + provider_config: CSVVectorOutputConfig = { "crs": pyproj.CRS.from_epsg(6372), "store": obstore.store.MemoryStore(), "results_prefix": results_prefix, @@ -390,17 +392,17 @@ def test_append_time_type_mismatch_timedelta_to_datetime(self, results_prefix): provider_config["overwrite"] = False csv_provider_append = CSVVectorOutputProvider(provider_config) - with pytest.raises(ValueError, match="time.*type|type.*mismatch"): + with pytest.raises(TypeError, match="time.*type|type.*mismatch"): csv_provider_append.write_vector( drainage_network, datetime(year=2020, month=3, day=23, hour=10) ) def test_append_time_type_mismatch_datetime_to_timedelta(self, results_prefix): - """Verify ValueError when appending timedelta to file with datetime.""" + """Verify TypeError when appending timedelta to file with datetime.""" drainage_network = create_dummy_drainage_network() # Write with datetime - provider_config = { + provider_config: CSVVectorOutputConfig = { "crs": pyproj.CRS.from_epsg(6372), "store": obstore.store.MemoryStore(), "results_prefix": results_prefix, @@ -414,7 +416,7 @@ def test_append_time_type_mismatch_datetime_to_timedelta(self, results_prefix): provider_config["overwrite"] = False csv_provider_append = CSVVectorOutputProvider(provider_config) - with pytest.raises(ValueError, match="time.*type|type.*mismatch"): + with pytest.raises(TypeError, match="time.*type|type.*mismatch"): csv_provider_append.write_vector(drainage_network, timedelta(seconds=60)) def test_append_node_ids_mismatch(self, results_prefix): @@ -423,7 +425,7 @@ def test_append_node_ids_mismatch(self, results_prefix): sim_time = timedelta(seconds=0) # Write initial data - provider_config = { + provider_config: CSVVectorOutputConfig = { "crs": pyproj.CRS.from_epsg(6372), "store": obstore.store.MemoryStore(), "results_prefix": results_prefix, @@ -476,7 +478,7 @@ def test_append_link_ids_mismatch(self, results_prefix): sim_time = timedelta(seconds=0) # Write initial data - provider_config = { + provider_config: CSVVectorOutputConfig = { "crs": pyproj.CRS.from_epsg(6372), "store": obstore.store.MemoryStore(), "results_prefix": results_prefix, @@ -520,7 +522,7 @@ def test_append_time_not_increasing(self, results_prefix): drainage_network = create_dummy_drainage_network() # Write initial data with two time steps - provider_config = { + provider_config: CSVVectorOutputConfig = { "crs": pyproj.CRS.from_epsg(6372), "store": obstore.store.MemoryStore(), "results_prefix": results_prefix, @@ -543,7 +545,7 @@ def test_append_time_equal_to_maximum(self, results_prefix): drainage_network = create_dummy_drainage_network() # Write initial data - provider_config = { + provider_config: CSVVectorOutputConfig = { "crs": pyproj.CRS.from_epsg(6372), "store": obstore.store.MemoryStore(), "results_prefix": results_prefix, diff --git a/tests/test_ea8a.py b/tests/test_ea8a.py index a975867..fd3d556 100644 --- a/tests/test_ea8a.py +++ b/tests/test_ea8a.py @@ -38,7 +38,7 @@ MemoryRasterOutputProvider, MemoryVectorOutputProvider, ) -from itzi_core.providers.xarray_input import XarrayRasterInputProvider +from itzi_core.providers.xarray_input import XarrayRasterInputConfig, XarrayRasterInputProvider from itzi_core.simulation_builder import SimulationBuilder # Mark all tests in this module as cloud tests @@ -84,11 +84,13 @@ def ea_test8a_xarray_data(test8a_file, test_data_path, test_data_temp_path): # Process DEM - import and aggregate from 50cm to 2m (matching GRASS r.resamp.stats) dem_path = os.path.join(unzip_path, "Test8DEM.asc") - dem_da = rioxarray.open_rasterio(dem_path, masked=True).isel(band=0) + dem_raster = rioxarray.open_rasterio(dem_path, masked=True) + assert isinstance(dem_raster, xr.DataArray) + dem_da = dem_raster.isel(band=0) # Aggregate using coarsen (4x4 blocks since 50cm to 2m is 4x factor) # This matches GRASS r.resamp.stats default (mean aggregation) - dem_da_coarse = dem_da.coarsen(x=4, y=4, boundary="pad").mean() + dem_da_coarse = dem_da.coarsen(x=4, y=4, boundary="pad").reduce(np.nanmean) # Interpolate to exact target coordinates dem_da_resampled = dem_da_coarse.interp(x=x_coords, y=y_coords, method="nearest") @@ -104,7 +106,9 @@ def ea_test8a_xarray_data(test8a_file, test_data_path, test_data_temp_path): # Process road pavement for Manning coefficient road_path = os.path.join(unzip_path, "Test8RoadPavement.asc") - road_da = rioxarray.open_rasterio(road_path, mask_and_scale=True).isel(band=0) + road_raster = rioxarray.open_rasterio(road_path, mask_and_scale=True) + assert isinstance(road_raster, xr.DataArray) + road_da = road_raster.isel(band=0) road_da = road_da.interp(x=x_coords, y=y_coords, method="nearest") road_data = road_da.values # Create Manning coefficient: 0.02 where road exists, 0.05 elsewhere @@ -270,7 +274,7 @@ def ea_test8a_sim(ea_test8a_xarray_data, test_data_path, test_data_temp_path): sim_end_time = sim_start_time + sim_duration # Create input provider - input_config = { + input_config: XarrayRasterInputConfig = { "dataset": ds, "input_map_names": { "dem": "dem", diff --git a/tests/test_flow.py b/tests/test_flow.py index 51ffcf3..49653ec 100644 --- a/tests/test_flow.py +++ b/tests/test_flow.py @@ -19,7 +19,6 @@ import numpy as np import pytest - from itzi_core.compute.partial_inertia_h import solve_h from itzi_core.compute.partial_inertia_q import solve_q @@ -262,16 +261,15 @@ def test_solve_h_optional_diagnostics(dtype, compute_vdir, compute_froude): expected_v = sqrt(10.0**2 + 6.0**2) expected_vdir = atan2(-6.0, 10.0) * 180.0 / pi % 360.0 expected_froude = expected_v / sqrt(9.81 * 0.1) - tolerance = {"rtol": 1e-6, "atol": 1e-6} - np.testing.assert_allclose(arr_v[1:-1, 1:-1], expected_v, **tolerance) - np.testing.assert_allclose(arr_vmax[1:-1, 1:-1], expected_v, **tolerance) + np.testing.assert_allclose(arr_v[1:-1, 1:-1], expected_v, rtol=1e-6, atol=1e-6) + np.testing.assert_allclose(arr_vmax[1:-1, 1:-1], expected_v, rtol=1e-6, atol=1e-6) if compute_vdir: - np.testing.assert_allclose(arr_vdir[1:-1, 1:-1], expected_vdir, **tolerance) + np.testing.assert_allclose(arr_vdir[1:-1, 1:-1], expected_vdir, rtol=1e-6, atol=1e-6) else: assert np.all(arr_vdir == dtype(-123.0)) if compute_froude: - np.testing.assert_allclose(arr_fr[1:-1, 1:-1], expected_froude, **tolerance) + np.testing.assert_allclose(arr_fr[1:-1, 1:-1], expected_froude, rtol=1e-6, atol=1e-6) else: assert np.all(arr_fr == dtype(-456.0)) diff --git a/tests/test_hotstart_state_loading.py b/tests/test_hotstart_state_loading.py index 50156a5..1dca73a 100644 --- a/tests/test_hotstart_state_loading.py +++ b/tests/test_hotstart_state_loading.py @@ -620,9 +620,7 @@ def test_build_rejects_drainage_mismatch_config_has_drainage( swmm_inp="fake.inp", # Has drainage ) - raster_output = MemoryRasterOutputProvider( - {"out_map_names": config_with_drainage.output_map_names} - ) + raster_output = MemoryRasterOutputProvider(config_with_drainage.output_map_names) with pytest.raises( HotstartError, match="Hotstart has no drainage state but current configuration" diff --git a/tests/test_hotstart_timed_inputs.py b/tests/test_hotstart_timed_inputs.py index effc727..a0c1c63 100644 --- a/tests/test_hotstart_timed_inputs.py +++ b/tests/test_hotstart_timed_inputs.py @@ -17,7 +17,7 @@ from collections.abc import Mapping, Sequence from datetime import datetime, timedelta -from typing import TYPE_CHECKING +from typing import TYPE_CHECKING, TypedDict import numpy as np import pytest @@ -46,13 +46,19 @@ } +class HotstartCheckpoint(TypedDict): + sim_time: datetime + rain: np.ndarray + water_depth: np.ndarray + + class MapAwareRasterInputProvider(RasterInputProvider): """Resolve canonical inputs through the current configured source names.""" def __init__( self, domain_data, - input_map_names: Mapping[str, str | None], + input_map_names: Mapping[str, str], source_slices: Mapping[str, Sequence[TimedRasterSlice]], start_time: datetime, end_time: datetime, @@ -242,13 +248,13 @@ def _build_provider_simulation( def _make_map_config( start_time: datetime, end_time: datetime, - input_map_names: dict[str, str | None], + input_map_names: Mapping[str, str], ) -> SimulationConfig: return _make_simulation_config( start_time, end_time, temporal_type=TemporalType.RELATIVE, - ).model_copy(update={"input_map_names": input_map_names}) + ).model_copy(update={"input_map_names": dict(input_map_names)}) def _source_slice( @@ -291,7 +297,7 @@ def _run_reference_with_hotstart_checkpoint( static_arrays: dict[str, np.ndarray], timed_arrays: dict[str, list[TimedRasterSlice]], split_target_time: datetime, -) -> tuple[dict[str, datetime | np.ndarray], bytes, "Simulation"]: +) -> tuple[HotstartCheckpoint, bytes, "Simulation"]: simulation = _build_provider_simulation( sim_config, domain_5by5, @@ -303,7 +309,7 @@ def _run_reference_with_hotstart_checkpoint( while simulation.sim_time < split_target_time: simulation.update() - checkpoint = { + checkpoint: HotstartCheckpoint = { "sim_time": simulation.sim_time, "rain": simulation.raster_domain.get_array("rain").copy(), "water_depth": simulation.raster_domain.get_array("water_depth").copy(), @@ -892,10 +898,9 @@ def test_non_cadence_aligned_end_writes_one_final_report(domain_5by5) -> None: simulation.update_until(end_time - start_time) simulation.finalize() - output_times = [ - sim_time - for sim_time, _ in simulation.report.raster_provider.output_maps_dict["water_depth"] - ] + raster_provider = simulation.report.raster_provider + assert isinstance(raster_provider, MemoryRasterOutputProvider) + output_times = [sim_time for sim_time, _ in raster_provider.output_maps_dict["water_depth"]] assert output_times == [ timedelta(seconds=0), timedelta(seconds=10), diff --git a/tests/test_icechunk_output.py b/tests/test_icechunk_output.py index 6ed1724..2588a90 100644 --- a/tests/test_icechunk_output.py +++ b/tests/test_icechunk_output.py @@ -30,7 +30,10 @@ import xarray as xr from itzi_core.array_definitions import ARRAY_DEFINITIONS, ArrayCategory -from itzi_core.providers.icechunk_output import IcechunkRasterOutputProvider +from itzi_core.providers.icechunk_output import ( + IcechunkRasterOutputConfig, + IcechunkRasterOutputProvider, +) # Mark all tests in this module as cloud tests pytestmark = pytest.mark.cloud @@ -75,10 +78,13 @@ def out_map_names(maps_dict: dict): @pytest.fixture def icechunk_provider( - temp_dir: tempfile.TemporaryDirectory, coordinates: dict, crs: pyproj.CRS, out_map_names: list + temp_dir: tempfile.TemporaryDirectory, + coordinates: dict, + crs: pyproj.CRS, + out_map_names: Mapping[str, str], ): storage = icechunk.local_filesystem_storage(temp_dir.name) - provider_config = { + provider_config: IcechunkRasterOutputConfig = { "out_map_names": out_map_names, "crs": crs, "x_coords": coordinates["x_coords"], @@ -344,7 +350,7 @@ def test_non_matching_shape( # Create and write original arrays (6x9 from fixture) storage = icechunk.local_filesystem_storage(temp_dir.name) - provider_config_1 = { + provider_config_1: IcechunkRasterOutputConfig = { "out_map_names": out_map_names, "crs": crs, "x_coords": coordinates["x_coords"], @@ -366,7 +372,7 @@ def test_non_matching_shape( new_coordinates = {"x_coords": new_x_coords, "y_coords": new_y_coords} # Create new provider with different dimensions but same storage - provider_config_2 = { + provider_config_2: IcechunkRasterOutputConfig = { "out_map_names": out_map_names, "crs": crs, "x_coords": new_coordinates["x_coords"], @@ -390,7 +396,7 @@ def test_non_matching_variable_names( # Create and write original arrays storage = icechunk.local_filesystem_storage(temp_dir.name) - provider_config_1 = { + provider_config_1: IcechunkRasterOutputConfig = { "out_map_names": out_map_names, "crs": crs, "x_coords": coordinates["x_coords"], @@ -405,7 +411,7 @@ def test_non_matching_variable_names( # Create provider with different variable names different_map_names = {key: value + "_different" for key, value in out_map_names.items()} - provider_config_2 = { + provider_config_2: IcechunkRasterOutputConfig = { "out_map_names": different_map_names, "crs": crs, "x_coords": coordinates["x_coords"], @@ -430,7 +436,7 @@ def test_non_matching_number_of_variables( # Create and write original arrays storage = icechunk.local_filesystem_storage(temp_dir.name) - provider_config_1 = { + provider_config_1: IcechunkRasterOutputConfig = { "out_map_names": out_map_names, "crs": crs, "x_coords": coordinates["x_coords"], @@ -446,7 +452,7 @@ def test_non_matching_number_of_variables( # Create provider with fewer variables fewer_map_names = dict(out_map_names.items()) del fewer_map_names["water_depth"] - provider_config_2 = { + provider_config_2: IcechunkRasterOutputConfig = { "out_map_names": fewer_map_names, "crs": crs, "x_coords": coordinates["x_coords"], @@ -471,7 +477,7 @@ def test_non_matching_coordinates_same_dimensions( # Create and write original arrays storage = icechunk.local_filesystem_storage(temp_dir.name) - provider_config_1 = { + provider_config_1: IcechunkRasterOutputConfig = { "out_map_names": out_map_names, "crs": crs, "x_coords": coordinates["x_coords"], @@ -489,7 +495,7 @@ def test_non_matching_coordinates_same_dimensions( different_y_coords = np.linspace(start=9999, stop=9999 + arr_shape[0], num=arr_shape[0]) different_x_coords = np.linspace(start=9999, stop=9999 + arr_shape[1], num=arr_shape[1]) - provider_config_2 = { + provider_config_2: IcechunkRasterOutputConfig = { "out_map_names": out_map_names, "crs": crs, "x_coords": different_x_coords, # Same shape, different values @@ -514,7 +520,7 @@ def test_non_matching_crs( # Create and write original arrays storage = icechunk.local_filesystem_storage(temp_dir.name) - provider_config_1 = { + provider_config_1: IcechunkRasterOutputConfig = { "out_map_names": out_map_names, "crs": crs, "x_coords": coordinates["x_coords"], @@ -529,7 +535,7 @@ def test_non_matching_crs( # Create provider with different CRS different_crs = pyproj.CRS.from_epsg(4326) # WGS84, different from Mexico LCC - provider_config_2 = { + provider_config_2: IcechunkRasterOutputConfig = { "out_map_names": out_map_names, "crs": different_crs, # Different CRS "x_coords": coordinates["x_coords"], @@ -563,7 +569,7 @@ def test_multi_session_data_persistence( # Session 1: Create first provider and write initial data storage = icechunk.local_filesystem_storage(temp_dir.name) - provider_config_1 = { + provider_config_1: IcechunkRasterOutputConfig = { "out_map_names": out_map_names, "crs": crs, "x_coords": coordinates["x_coords"], @@ -579,7 +585,7 @@ def test_multi_session_data_persistence( icechunk_p1.write_arrays(maps_dict_session1, sim_time_2) # Write same data twice # Session 2: Create new provider instance with same storage and compatible config - provider_config_2 = { + provider_config_2: IcechunkRasterOutputConfig = { "out_map_names": out_map_names, "crs": crs, "x_coords": coordinates["x_coords"], diff --git a/tests/test_mass_balance_provider.py b/tests/test_mass_balance_provider.py index 6a4f68c..074cebf 100644 --- a/tests/test_mass_balance_provider.py +++ b/tests/test_mass_balance_provider.py @@ -122,15 +122,17 @@ def test_simulation_config_rejects_removed_stats_file() -> None: start_time = datetime(2000, 1, 1, tzinfo=UTC) with pytest.raises(ValidationError, match="stats_file"): - SimulationConfig( - start_time=start_time, - end_time=start_time + timedelta(seconds=1), - record_step=timedelta(seconds=1), - temporal_type=TemporalType.RELATIVE, - input_map_names={}, - output_map_names={}, - surface_flow_parameters=SurfaceFlowParameters(), - stats_file="removed.csv", + SimulationConfig.model_validate( + { + "start_time": start_time, + "end_time": start_time + timedelta(seconds=1), + "record_step": timedelta(seconds=1), + "temporal_type": TemporalType.RELATIVE, + "input_map_names": {}, + "output_map_names": {}, + "surface_flow_parameters": SurfaceFlowParameters(), + "stats_file": "removed.csv", + } ) diff --git a/tests/test_memory_input.py b/tests/test_memory_input.py index 0347a94..5183e33 100644 --- a/tests/test_memory_input.py +++ b/tests/test_memory_input.py @@ -15,12 +15,17 @@ from __future__ import annotations from datetime import datetime, timedelta +from typing import Any, cast import numpy as np import pytest from itzi_core.providers.domain_data import DomainData -from itzi_core.providers.memory_input import MemoryRasterInputProvider, TimedRasterSlice +from itzi_core.providers.memory_input import ( + MemoryRasterInputConfig, + MemoryRasterInputProvider, + TimedRasterSlice, +) @pytest.fixture @@ -46,15 +51,15 @@ def simulation_times() -> dict[str, datetime]: def make_config( domain_data: DomainData, simulation_times: dict[str, datetime], - **overrides, -) -> dict: - config = { + **overrides: Any, +) -> MemoryRasterInputConfig: + config: dict[str, Any] = { "domain_data": domain_data, "simulation_start_time": simulation_times["start_time"], "simulation_end_time": simulation_times["end_time"], } config.update(overrides) - return config + return cast(MemoryRasterInputConfig, config) def test_provider_creation_with_empty_arrays( diff --git a/tests/test_swmm_hotstart_roundtrip.py b/tests/test_swmm_hotstart_roundtrip.py index 24b4369..86f9433 100644 --- a/tests/test_swmm_hotstart_roundtrip.py +++ b/tests/test_swmm_hotstart_roundtrip.py @@ -16,22 +16,21 @@ from __future__ import annotations -from datetime import datetime, timedelta -from pathlib import Path import os import shutil import tempfile -from typing import Any +from datetime import datetime, timedelta +from pathlib import Path +from typing import Any, TypedDict import numpy as np -import pytest import pyswmm +import pytest from pyswmm.simulation import _sim_state_instance from pyswmm.toolkitapi import NodeResults, SimulationTime from itzi_core import SwmmInputParser - SECONDS_PER_DAY = 24 * 3600 SPLIT_TIME = timedelta(hours=1, minutes=40) SWMM_J1_TOTAL_INFLOW_MAX_DIFF = 0.25 @@ -39,6 +38,34 @@ SWMM_J1_INTEGRATED_INFLOW_MAX_DIFF = 2.0 +class SnapshotAccumulator(TypedDict): + elapsed_seconds: list[float] + node_depths: list[list[float]] + node_heads: list[list[float]] + node_total_inflow: list[list[float]] + node_cumulative_inflow: list[list[float]] + node_volumes: list[list[float]] + node_overflow: list[list[float]] + link_flows: list[list[float]] + link_depths: list[list[float]] + link_volumes: list[list[float]] + + +class SnapshotResult(TypedDict): + node_ids: tuple[str, ...] + link_ids: tuple[str, ...] + elapsed_seconds: np.ndarray + node_depths: np.ndarray + node_heads: np.ndarray + node_total_inflow: np.ndarray + node_cumulative_inflow: np.ndarray + node_volumes: np.ndarray + node_overflow: np.ndarray + link_flows: np.ndarray + link_depths: np.ndarray + link_volumes: np.ndarray + + @pytest.fixture def ea8b_inp_path(test_data_path: str, tmp_path: Path) -> Path: source_inp = Path(test_data_path) / "EA_test_8" / "b" / "test8b_drainage_ponding.inp" @@ -48,7 +75,7 @@ def ea8b_inp_path(test_data_path: str, tmp_path: Path) -> Path: def _record_snapshot( - snapshots: dict[str, list[float] | list[list[float]]], + snapshots: SnapshotAccumulator, node_objects: list, link_objects: list, swmm_model: Any, @@ -69,7 +96,7 @@ def _record_snapshot( snapshots["link_volumes"].append([link.volume for link in link_objects]) -def _empty_snapshots() -> dict[str, list[float] | list[list[float]]]: +def _empty_snapshots() -> SnapshotAccumulator: return { "elapsed_seconds": [], "node_depths": [], @@ -85,10 +112,10 @@ def _empty_snapshots() -> dict[str, list[float] | list[list[float]]]: def _finalize_snapshots( - snapshots: dict[str, list[float] | list[list[float]]], + snapshots: SnapshotAccumulator, node_ids: tuple[str, ...], link_ids: tuple[str, ...], -) -> dict[str, tuple[str, ...] | np.ndarray]: +) -> SnapshotResult: return { "node_ids": node_ids, "link_ids": link_ids, @@ -232,9 +259,7 @@ def _run_ab_ponding_diagnostic( return results -def _run_uninterrupted( - inp_file: str, split_seconds: float -) -> dict[str, tuple[str, ...] | np.ndarray]: +def _run_uninterrupted(inp_file: str, split_seconds: float) -> SnapshotResult: swmm_sim = pyswmm.Simulation(inp_file) swmm_model = swmm_sim._model node_objects = list(pyswmm.Nodes(swmm_sim)) @@ -276,9 +301,7 @@ def _run_uninterrupted( return _finalize_snapshots(snapshots, node_ids, link_ids) -def _run_with_hotstart( - inp_file: str, split_seconds: float -) -> dict[str, tuple[str, ...] | np.ndarray]: +def _run_with_hotstart(inp_file: str, split_seconds: float) -> SnapshotResult: parser = SwmmInputParser(inp_file) original_start = parser.get_start_datetime() assert original_start is not None, "Failed to parse SWMM START_DATE/START_TIME" diff --git a/tests/test_xarray_input.py b/tests/test_xarray_input.py index e23818b..1710af0 100644 --- a/tests/test_xarray_input.py +++ b/tests/test_xarray_input.py @@ -12,8 +12,8 @@ GNU Lesser General Public License for more details. """ -from typing import Dict from datetime import datetime, timedelta +from typing import Dict import numpy as np import pandas as pd @@ -23,12 +23,11 @@ pytest.importorskip("xarray") pytest.importorskip("pyproj") -import xarray as xr import pyproj +import xarray as xr -from itzi_core.providers.xarray_input import XarrayRasterInputProvider, XarrayRasterInputConfig from itzi_core.const import TemporalType - +from itzi_core.providers.xarray_input import XarrayRasterInputConfig, XarrayRasterInputProvider # Mark all tests in this module as cloud tests pytestmark = pytest.mark.cloud @@ -333,6 +332,7 @@ def test_xarray_input_provider_uses_half_open_windows_at_exact_boundary( current_time = datetime(2023, 1, 1, 2, 0, 0) array, start_time, end_time = provider.get_array("rainfall", current_time) + assert array is not None expected_time_index = 2 base_data = xarray_input_data["input_maps_dict"]["rainfall"] @@ -358,6 +358,7 @@ def test_xarray_input_provider_extends_last_slice_to_simulation_end( current_time = datetime(2023, 1, 1, 4, 30, 0) array, start_time, end_time = provider.get_array("rainfall", current_time) + assert array is not None expected_time_index = 4 base_data = xarray_input_data["input_maps_dict"]["rainfall"] @@ -859,7 +860,9 @@ def mixed_dimensions_data(input_maps_dict: Dict, coordinates: Dict, crs: pyproj. time_step_hours: int = 1 # Relative time coordinates (timedelta) - relative_times: list[int] = [timedelta(hours=i * time_step_hours) for i in range(time_steps)] + relative_times: list[timedelta] = [ + timedelta(hours=i * time_step_hours) for i in range(time_steps) + ] # Absolute time coordinates (datetime) start_time = datetime(2023, 1, 1, 0, 0, 0) @@ -871,18 +874,18 @@ def mixed_dimensions_data(input_maps_dict: Dict, coordinates: Dict, crs: pyproj. data_vars = {} # 1. 2D static array with (lat, lon) dimensions - dem_data: str = input_maps_dict["dem"] + dem_data: np.ndarray = input_maps_dict["dem"] data_vars["elevation"] = (["lat", "lon"], dem_data) # 2. 3D array with relative time (rel_time, northing, easting) - rainfall_data: str = input_maps_dict["rainfall"] + rainfall_data: np.ndarray = input_maps_dict["rainfall"] rainfall_time_data: np.ndarray = np.stack( [rainfall_data * (1 + 0.1 * t) for t in range(time_steps)] ) data_vars["precip"] = (["rel_time", "northing", "easting"], rainfall_time_data) # 3. 3D array with absolute time (abs_time, rows, cols) - bc_data: str = input_maps_dict["boundary_conditions"] + bc_data: np.ndarray = input_maps_dict["boundary_conditions"] bc_time_data: np.ndarray = np.stack([bc_data * (1 + 0.2 * t) for t in range(time_steps)]) data_vars["boundary"] = (["abs_time", "rows", "cols"], bc_time_data)