Skip to content

Commit ee50e8b

Browse files
committed
TYP: Adjust ERT typehints to new import locations
1 parent d3ac331 commit ee50e8b

6 files changed

Lines changed: 33 additions & 26 deletions

File tree

src/fmu/dataio/_workflows/case/_config.py

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -13,7 +13,7 @@
1313
from fmu.settings import ProjectFMUDirectory
1414

1515
if TYPE_CHECKING:
16-
import ert
16+
from ert.runpaths import Runpaths as ErtRunpaths
1717

1818
logger: Final = logging.getLogger(__name__)
1919
logger.setLevel(logging.CRITICAL)
@@ -53,7 +53,7 @@ def validate(self) -> None:
5353
@classmethod
5454
def from_presim_workflow(
5555
cls,
56-
run_paths: ert.Runpaths,
56+
run_paths: ErtRunpaths,
5757
args: argparse.Namespace,
5858
fmu_dir: ProjectFMUDirectory | None = None,
5959
) -> Self:

src/fmu/dataio/_workflows/case/_observations.py

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -9,7 +9,7 @@
99
import pyarrow as pa
1010

1111
if TYPE_CHECKING:
12-
import ert
12+
from ert.storage import Ensemble as ErtEnsemble
1313

1414

1515
logger: Final = logging.getLogger(__name__)
@@ -46,7 +46,7 @@ def _prepare_observations_dataframe(
4646

4747

4848
def get_ert_observations_table(
49-
ensemble: ert.Ensemble, obs_type: Literal["rft", "summary", "breakthrough"]
49+
ensemble: ErtEnsemble, obs_type: Literal["rft", "summary", "breakthrough"]
5050
) -> pa.Table | None:
5151
"""Extract observations from ert storage and process it into an arrow table."""
5252
logger.info(f"Observation type: {obs_type}")

src/fmu/dataio/_workflows/case/_parameters.py

Lines changed: 3 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -13,6 +13,7 @@
1313

1414
if TYPE_CHECKING:
1515
import polars as pl
16+
from ert.storage import Ensemble as ErtEnsemble
1617

1718

1819
logger: Final = logging.getLogger(__name__)
@@ -51,7 +52,7 @@ def _resolve_pa_field_type(name: str, pa_type: pa.DataType) -> pa.DataType:
5152

5253

5354
def _process_parameters(
54-
scalars_df: pl.DataFrame, ensemble: ert.Ensemble
55+
scalars_df: pl.DataFrame, ensemble: ErtEnsemble
5556
) -> tuple[pa.Table, list[int]]:
5657
"""Process parameters into an Arrow table with metadata."""
5758
import pyarrow as pa
@@ -96,7 +97,7 @@ def _process_parameters(
9697
return table, realizations
9798

9899

99-
def get_ert_parameters_table(ensemble: ert.Ensemble) -> pa.Table | None:
100+
def get_ert_parameters_table(ensemble: ErtEnsemble) -> pa.Table | None:
100101
"""Exports Ert parameters as a Parquet file as the ensemble level."""
101102

102103
scalars_df = ensemble.load_scalars()

src/fmu/dataio/_workflows/case/main.py

Lines changed: 19 additions & 14 deletions
Original file line numberDiff line numberDiff line change
@@ -9,7 +9,7 @@
99
import logging
1010
import shutil
1111
from pathlib import Path
12-
from typing import Final
12+
from typing import TYPE_CHECKING, Final
1313

1414
import ert
1515

@@ -34,6 +34,11 @@
3434
from ._parameters import get_ert_parameters_table
3535
from .export_case_metadata import ExportCaseMetadata
3636

37+
if TYPE_CHECKING:
38+
from ert.runpaths import Runpaths as ErtRunpaths
39+
from ert.storage import Ensemble as ErtEnsemble
40+
41+
3742
logger: Final = logging.getLogger(__name__)
3843
logger.setLevel(logging.CRITICAL)
3944

@@ -58,8 +63,8 @@
5863

5964

6065
def _get_ensemble_name(
61-
ensemble: ert.Ensemble,
62-
run_paths: ert.Runpaths,
66+
ensemble: ErtEnsemble,
67+
run_paths: ErtRunpaths,
6368
casepath: Path,
6469
) -> str:
6570
"""Determine ensemble name from run path.
@@ -77,7 +82,7 @@ def _get_ensemble_name(
7782

7883

7984
def _queue_ert_parameters(
80-
ensemble: ert.Ensemble,
85+
ensemble: ErtEnsemble,
8186
ensemble_name: str,
8287
workflow_config: CaseWorkflowConfig,
8388
sumo_uploader: SumoUploaderInterface,
@@ -107,7 +112,7 @@ def _queue_ert_parameters(
107112

108113

109114
def _queue_ert_observations_breakthrough(
110-
ensemble: ert.Ensemble,
115+
ensemble: ErtEnsemble,
111116
ensemble_name: str,
112117
workflow_config: CaseWorkflowConfig,
113118
sumo_uploader: SumoUploaderInterface,
@@ -139,7 +144,7 @@ def _queue_ert_observations_breakthrough(
139144

140145

141146
def _queue_ert_observations_rft(
142-
ensemble: ert.Ensemble,
147+
ensemble: ErtEnsemble,
143148
ensemble_name: str,
144149
workflow_config: CaseWorkflowConfig,
145150
sumo_uploader: SumoUploaderInterface,
@@ -170,7 +175,7 @@ def _queue_ert_observations_rft(
170175

171176

172177
def _queue_ert_observations_summary(
173-
ensemble: ert.Ensemble,
178+
ensemble: ErtEnsemble,
174179
ensemble_name: str,
175180
workflow_config: CaseWorkflowConfig,
176181
sumo_uploader: SumoUploaderInterface,
@@ -233,8 +238,8 @@ def _queue_stratigraphy_mappings(
233238

234239

235240
def _upload_files_to_sumo(
236-
ensemble: ert.Ensemble,
237-
run_paths: ert.Runpaths,
241+
ensemble: ErtEnsemble,
242+
run_paths: ErtRunpaths,
238243
workflow_config: CaseWorkflowConfig,
239244
sumo_uploader: SumoUploaderInterface,
240245
) -> None:
@@ -256,8 +261,8 @@ def _upload_files_to_sumo(
256261

257262

258263
def _run_workflow(
259-
ensemble: ert.Ensemble,
260-
run_paths: ert.Runpaths,
264+
ensemble: ErtEnsemble,
265+
run_paths: ErtRunpaths,
261266
workflow_config: CaseWorkflowConfig,
262267
) -> None:
263268
"""Main workflow entry point."""
@@ -362,8 +367,8 @@ class WfExportCaseMetadata(ert.ErtScript):
362367
def run(
363368
self,
364369
workflow_args: list[str],
365-
ensemble: ert.Ensemble,
366-
run_paths: ert.Runpaths,
370+
ensemble: ErtEnsemble,
371+
run_paths: ErtRunpaths,
367372
) -> None:
368373
"""Parse arguments and run the workflow."""
369374
parser = get_parser()
@@ -376,7 +381,7 @@ def run(
376381

377382

378383
@ert.plugin(name="fmu_dataio")
379-
def ertscript_workflow(config: ert.CaseWorkflowConfigs) -> None:
384+
def ertscript_workflow(config: ert.WorkflowConfigs) -> None:
380385
"""Hook the WfExportCaseMetadata class with documentation into ERT."""
381386
config.add_workflow(
382387
WfExportCaseMetadata,

tests/test_ert_integration/test_wf_create_case_metadata.py

Lines changed: 5 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -53,6 +53,7 @@
5353
)
5454

5555
if TYPE_CHECKING:
56+
from ert.storage import Ensemble as ErtEnsemble
5657
from fmu.datamodels.fmu_results.global_configuration import GlobalConfiguration
5758

5859

@@ -462,7 +463,7 @@ def test_create_case_metadata_collects_ert_parameters_as_expected(
462463
scalars_and_config = []
463464

464465
def capture_params(
465-
ensemble: ert.Ensemble,
466+
ensemble: ErtEnsemble,
466467
ensemble_name: str,
467468
workflow_config: CaseWorkflowConfig,
468469
sumo_uploader: SumoUploaderInterface,
@@ -797,7 +798,7 @@ def test_create_case_metadata_collects_rft_observations_as_expected(
797798
captured_tables = {}
798799

799800
def capture_observation_tables(
800-
ensemble: ert.Ensemble,
801+
ensemble: ErtEnsemble,
801802
obs_type: str,
802803
) -> None:
803804
"""Captures observation tables from Ert run.
@@ -808,7 +809,7 @@ def capture_observation_tables(
808809
return df
809810

810811
def mock_create_observation_dataframes(
811-
observations: ert.Ensemble,
812+
observations: ErtEnsemble,
812813
shape_registry: ShapeRegistry,
813814
) -> dict[str, pl.DataFrame]:
814815
"""mock"""
@@ -868,7 +869,7 @@ def test_create_case_metadata_with_no_observations(
868869
captured_tables = {}
869870

870871
def capture_observation_tables(
871-
ensemble: ert.Ensemble,
872+
ensemble: ErtEnsemble,
872873
obs_type: str,
873874
) -> None:
874875
"""Captures rft observations from Ert run"""

tests/test_units/test_workflows/test_case_workflow.py

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -22,7 +22,7 @@
2222

2323
@pytest.fixture
2424
def mock_ensemble() -> Callable[[int], MagicMock]:
25-
"""Mocks an ert.Ensemble object."""
25+
"""Mocks an ert.storage.Ensemble object."""
2626

2727
def _mock_ensemble(iteration: int = 0) -> MagicMock:
2828
"""Creates the mocked object."""
@@ -35,7 +35,7 @@ def _mock_ensemble(iteration: int = 0) -> MagicMock:
3535

3636
@pytest.fixture
3737
def mock_run_paths() -> Callable[[str], MagicMock]:
38-
"""Mocks and ert.Runpaths object."""
38+
"""Mocks and ert.runpaths.Runpaths object."""
3939

4040
def _mock_run_paths(runpath: str = "/tmp/realization-0/iter-0") -> MagicMock:
4141
"""Creates the mocked object."""

0 commit comments

Comments
 (0)