From a5a366ad5b55b0e489942830dc80c0fa50915be4 Mon Sep 17 00:00:00 2001 From: Benjamin A Chandler Date: Mon, 24 Aug 2026 14:11:31 +0200 Subject: [PATCH] Update PlotApi to use StorageApi for fetching data --- src/ert/gui/plotting/plot_api.py | 640 ++++++++---------- src/ert/gui/plotting/plot_window.py | 11 +- src/ert/services/ert_client.py | 21 +- src/ert/services/shared_client.py | 7 + .../test_dark_storage_performance.py | 20 +- tests/ert/ui_tests/gui/conftest.py | 8 + .../ert/unit_tests/gui/tools/plot/conftest.py | 17 +- .../gui/tools/plot/test_plot_api.py | 72 +- .../gui/tools/plot/test_plot_window.py | 20 +- .../unit_tests/services/test_ert_client.py | 4 +- 10 files changed, 362 insertions(+), 458 deletions(-) diff --git a/src/ert/gui/plotting/plot_api.py b/src/ert/gui/plotting/plot_api.py index 7819ec6bc7d..3f592a9faa9 100644 --- a/src/ert/gui/plotting/plot_api.py +++ b/src/ert/gui/plotting/plot_api.py @@ -1,13 +1,10 @@ from __future__ import annotations -import io -import json import logging from dataclasses import dataclass from functools import cached_property, lru_cache from itertools import combinations as combi from typing import TYPE_CHECKING, Any, NamedTuple -from urllib.parse import quote import httpx import numpy as np @@ -22,7 +19,7 @@ KnownResponseTypes, ) from ert.config.response_config import ResponseConfig -from ert.services import create_ertserver_client +from ert.services.ert_client import ErtClient from ert.storage.local_experiment import _parameters_adapter as parameter_config_adapter from ert.storage.local_experiment import _responses_adapter as response_config_adapter from ert.storage.realization_storage_state import RealizationStorageState @@ -34,9 +31,6 @@ from pathlib import Path -TIMEOUT = 120 - - @dataclass(frozen=True, eq=True) class EnsembleObject: name: str @@ -60,26 +54,26 @@ class PlotApiKeyDefinition(NamedTuple): class PlotApi: - def __init__(self, ens_path: Path) -> None: + def __init__(self, ens_path: Path, ert_client: ErtClient | None = None) -> None: self.ens_path: Path = ens_path + self._client = ert_client or ErtClient.get_client(ens_path) self._all_ensembles: list[EnsembleObject] | None = None + # Caches are bound per instance so they are dropped along with the + # PlotApi, and so that key normalization happens before the lookup. + self._cached_gradient = process_arg( + key="key", process=lambda s: s.split("@", maxsplit=1)[0] + )(lru_cache(maxsize=256)(self._fetch_gradient)) + self._cached_parameter = lru_cache(maxsize=256)(self._fetch_parameter) + self._cached_controls = lru_cache(maxsize=128)(self._fetch_controls) + @property def api_version(self) -> str: - with create_ertserver_client(self.ens_path) as client: - try: - http_response = client.get("/version", timeout=TIMEOUT) - self._check_http_response(http_response) - api_version = str(http_response.json()) - except Exception as exc: - logger.exception(exc) - raise exc - else: - return api_version - - @staticmethod - def escape(s: str) -> str: - return quote(quote(s, safe="")) + try: + return self._client.version() + except Exception as exc: + logger.exception(exc) + raise def _get_ensemble_by_id(self, id_: str) -> EnsembleObject | None: for ensemble in self.get_all_ensembles(): @@ -92,90 +86,61 @@ def get_all_ensembles(self) -> list[EnsembleObject]: return self._all_ensembles self._all_ensembles = [] - with create_ertserver_client(self.ens_path) as client: - try: # ruff: ignore[too-many-statements-in-try-clause] - http_response = client.get("/experiments", timeout=TIMEOUT) - self._check_http_response(http_response) - experiments = http_response.json() - for experiment in experiments: - for ensemble_id in experiment["ensemble_ids"]: - http_response = client.get( - f"/ensembles/{ensemble_id}", timeout=TIMEOUT + try: # ruff: ignore[too-many-statements-in-try-clause] + for experiment in self._client.experiments(): + for ensemble_id in experiment["ensemble_ids"]: + response_json = self._client.ensemble(ensemble_id) + ensemble_name: str = response_json["userdata"]["name"] + experiment_name: str = response_json["userdata"]["experiment_name"] + ensemble_started_at = response_json["userdata"]["started_at"] + ensemble_undefined = False + if realization_storage_states := response_json.get( + "realization_storage_states" + ): + ensemble_undefined = ( + RealizationStorageState.PARAMETERS_LOADED + not in set(realization_storage_states) ) - self._check_http_response(http_response) - response_json: dict[str, Any] = http_response.json() - ensemble_name: str = response_json["userdata"]["name"] - experiment_name: str = response_json["userdata"][ - "experiment_name" - ] - ensemble_started_at = response_json["userdata"]["started_at"] - ensemble_undefined = False - if realization_storage_states := response_json.get( - "realization_storage_states" - ): - ensemble_undefined = ( - RealizationStorageState.PARAMETERS_LOADED - not in set(realization_storage_states) - ) - self._all_ensembles.append( - EnsembleObject( - name=ensemble_name, - id=ensemble_id, - experiment_name=experiment_name, - hidden=ensemble_name.startswith(".") - or ensemble_undefined, - started_at=ensemble_started_at, - has_func_eval=bool( - response_json["userdata"].get( - "has_func_eval", False - ) - ), - has_gradient=bool( - response_json["userdata"].get("has_gradient", False) - ), - ) + self._all_ensembles.append( + EnsembleObject( + name=ensemble_name, + id=ensemble_id, + experiment_name=experiment_name, + hidden=ensemble_name.startswith(".") or ensemble_undefined, + started_at=ensemble_started_at, + has_func_eval=bool( + response_json["userdata"].get("has_func_eval", False) + ), + has_gradient=bool( + response_json["userdata"].get("has_gradient", False) + ), ) - except IndexError as exc: - logger.exception(exc) - raise exc - else: - return self._all_ensembles - - @staticmethod - def _check_http_response(http_response: httpx._models.Response) -> None: - if http_response.status_code == httpx.codes.UNAUTHORIZED: - raise httpx.RequestError(message=f"{http_response.text}") - if http_response.status_code != httpx.codes.OK: - raise httpx.RequestError( - f" Please report this error and try restarting the application." - f"{http_response.text} from url: {http_response.url}." - ) + ) + except IndexError as exc: + logger.exception(exc) + raise + else: + return self._all_ensembles @cached_property def parameters_api_key_defs(self) -> list[PlotApiKeyDefinition]: all_keys: dict[str, PlotApiKeyDefinition] = {} - all_params = {} - - with create_ertserver_client(self.ens_path) as client: - http_response = client.get("/experiments", timeout=TIMEOUT) - self._check_http_response(http_response) - for experiment in http_response.json(): - for metadata in experiment["parameters"].values(): - param_cfg = parameter_config_adapter.validate_python(metadata) - if group := metadata.get("group"): - param_key = f"{group}:{metadata['name']}" - else: - param_key = metadata["name"] - all_keys[param_key] = PlotApiKeyDefinition( - key=param_key, - index_type=None, - observations=False, - dimensionality=metadata["dimensionality"], - metadata={"data_origin": metadata["type"]}, - parameter=param_cfg, - ) - all_params[param_key] = all_keys[param_key] + for experiment in self._client.experiments(): + for metadata in experiment["parameters"].values(): + param_cfg = parameter_config_adapter.validate_python(metadata) + if group := metadata.get("group"): + param_key = f"{group}:{metadata['name']}" + else: + param_key = metadata["name"] + all_keys[param_key] = PlotApiKeyDefinition( + key=param_key, + index_type=None, + observations=False, + dimensionality=metadata["dimensionality"], + metadata={"data_origin": metadata["type"]}, + parameter=param_cfg, + ) return list(all_keys.values()) @@ -183,72 +148,68 @@ def parameters_api_key_defs(self) -> list[PlotApiKeyDefinition]: def responses_api_key_defs(self) -> list[PlotApiKeyDefinition]: key_defs: dict[str, PlotApiKeyDefinition] = {} - with create_ertserver_client(self.ens_path) as client: - http_response = client.get("/experiments", timeout=TIMEOUT) - self._check_http_response(http_response) - - def update_keydef(plot_key_def: PlotApiKeyDefinition) -> None: - # Only replace existing key definition if the new has observations - if plot_key_def.key not in key_defs or plot_key_def.observations: - key_defs[plot_key_def.key] = plot_key_def + def update_keydef(plot_key_def: PlotApiKeyDefinition) -> None: + # Only replace existing key definition if the new has observations + if plot_key_def.key not in key_defs or plot_key_def.observations: + key_defs[plot_key_def.key] = plot_key_def - for experiment in http_response.json(): - for response_type, metadata in experiment["responses"].items(): - response_config: KnownResponseTypes = ( - response_config_adapter.validate_python(metadata) + for experiment in self._client.experiments(): + for response_type, metadata in experiment["responses"].items(): + response_config: KnownResponseTypes = ( + response_config_adapter.validate_python(metadata) + ) + keys = response_config.response_keys() + for key in keys: + has_obs = ( + response_type in experiment["observations"] + and key in experiment["observations"][response_type] ) - keys = response_config.response_keys() - for key in keys: - has_obs = ( - response_type in experiment["observations"] - and key in experiment["observations"][response_type] - ) - if response_config.filter_on is not None: - # Only assume one filter_on, this code is to be - # considered a bit "temp". - # In general, we could create a dropdown per - # filter_on on the frontend side - - filter_for_key = response_config.filter_on.get(key, {}) - for filter_key, values in filter_for_key.items(): - for v in values: - filter_on = {filter_key: v} - subkey = f"{key}@{v}" - update_keydef( - PlotApiKeyDefinition( - key=subkey, - index_type="VALUE", - observations=has_obs, - dimensionality=2, - metadata={ - "data_origin": response_type, - }, - filter_on=filter_on, - response=response_config, - ) + if response_config.filter_on is not None: + # Only assume one filter_on, this code is to be + # considered a bit "temp". + # In general, we could create a dropdown per + # filter_on on the frontend side + + filter_for_key = response_config.filter_on.get(key, {}) + for filter_key, values in filter_for_key.items(): + for v in values: + filter_on = {filter_key: v} + subkey = f"{key}@{v}" + update_keydef( + PlotApiKeyDefinition( + key=subkey, + index_type="VALUE", + observations=has_obs, + dimensionality=2, + metadata={ + "data_origin": response_type, + }, + filter_on=filter_on, + response=response_config, ) - else: - update_keydef( - PlotApiKeyDefinition( - key=key, - index_type="VALUE", - observations=has_obs, - dimensionality=2, - metadata={"data_origin": response_type}, - response=response_config, ) + else: + update_keydef( + PlotApiKeyDefinition( + key=key, + index_type="VALUE", + observations=has_obs, + dimensionality=2, + metadata={"data_origin": response_type}, + response=response_config, ) - - if "everest_objectives" in experiment["responses"]: - update_keydef( - PlotApiKeyDefinition( - key="total objective value", - index_type="VALUE", - observations=False, - dimensionality=2, - metadata={"data_origin": "everest_batch_objectives"}, ) + + if "everest_objectives" in experiment["responses"]: + update_keydef( + PlotApiKeyDefinition( + key="total objective value", + index_type="VALUE", + observations=False, + dimensionality=2, + metadata={"data_origin": "everest_batch_objectives"}, ) + ) return list(key_defs.values()) @@ -268,117 +229,100 @@ def data_for_response( if "@" in response_key: response_key = response_key.split("@", maxsplit=1)[0] - with create_ertserver_client(self.ens_path) as client: - http_response = client.get( - f"/ensembles/{ensemble_id}/responses/{PlotApi.escape(response_key)}", - headers={"accept": "application/x-parquet"}, - params={"filter_on": json.dumps(filter_on)} - if filter_on is not None - else None, - timeout=TIMEOUT, - ) - self._check_http_response(http_response) - - stream = io.BytesIO(http_response.content) - df = pd.read_parquet(stream) - if df.empty: - return df + df = self._client.ert_response(ensemble_id, response_key, filter_on) - if is_everest: - assert {"batch_id", "realization"}.issubset(df.columns) + if df.empty: + return df - float_columns = [ - col for col in df.columns if col not in {"batch_id", "realization"} - ] + if is_everest: + assert {"batch_id", "realization"}.issubset(df.columns) - return df.astype( - dict.fromkeys(float_columns, float) - | { - "batch_id": int, - "realization": int, - } - ) + float_columns = [ + col for col in df.columns if col not in {"batch_id", "realization"} + ] - if ( - key_def is not None - and key_def.metadata.get("data_origin") == "everest_batch_objectives" - ): - assert {"batch_id", "accepted"}.issubset(df.columns) + return df.astype( + dict.fromkeys(float_columns, float) + | { + "batch_id": int, + "realization": int, + } + ) - float_columns_names = ( - {"batch_id", "accepted", "constraint_violation_type"} - if "constraint_violation_type" in df.columns - else {"batch_id", "accepted"} - ) - float_columns = [ - col for col in df.columns if col not in float_columns_names - ] + if ( + key_def is not None + and key_def.metadata.get("data_origin") == "everest_batch_objectives" + ): + assert {"batch_id", "accepted"}.issubset(df.columns) - return df.astype( - dict.fromkeys(float_columns, float) - | { - "batch_id": int, - "accepted": bool, - "improvement_value": float, - "constraint_violation_type": str, - } - if "constraint_violation_type" in df.columns - else { - "batch_id": int, - "accepted": bool, - "improvement_value": float, - } - ) + float_columns_names = ( + {"batch_id", "accepted", "constraint_violation_type"} + if "constraint_violation_type" in df.columns + else {"batch_id", "accepted"} + ) + float_columns = [ + col for col in df.columns if col not in float_columns_names + ] - try: - df.columns = pd.to_datetime(df.columns, format="%Y-%m-%d %H:%M:%S") - except (ParserError, ValueError): - try: - df.columns = [int(s) for s in df.columns] - except ValueError: - df.columns = [float(s) for s in df.columns] + return df.astype( + dict.fromkeys(float_columns, float) + | { + "batch_id": int, + "accepted": bool, + "improvement_value": float, + "constraint_violation_type": str, + } + if "constraint_violation_type" in df.columns + else { + "batch_id": int, + "accepted": bool, + "improvement_value": float, + } + ) + try: + df.columns = pd.to_datetime(df.columns, format="%Y-%m-%d %H:%M:%S") + except (ParserError, ValueError): try: - return df.astype(float) + df.columns = [int(s) for s in df.columns] except ValueError: - return df - - @staticmethod - @process_arg(key="key", process=lambda s: s.split("@", maxsplit=1)[0]) - @lru_cache(maxsize=256) - def data_for_gradient(ensemble_id: str, key: str, ens_path: Path) -> pd.DataFrame: - with create_ertserver_client(ens_path) as client: - http_response = client.get( - f"/ensembles/{ensemble_id}/gradients/{PlotApi.escape(key)}", - headers={"accept": "application/x-parquet"}, - timeout=TIMEOUT, - ) - PlotApi._check_http_response(http_response) + df.columns = [float(s) for s in df.columns] - stream = io.BytesIO(http_response.content) - df = pd.read_parquet(stream) + try: + return df.astype(float) + except ValueError: + return df - if df.empty: - return df + def data_for_gradient(self, ensemble_id: str, key: str) -> pd.DataFrame: + return self._cached_gradient(ensemble_id, key) - return df.astype( - { - "batch_id": int, - "control_name": str, - key: float, - } - ) + def _fetch_gradient(self, ensemble_id: str, key: str) -> pd.DataFrame: + df = self._client.gradient(ensemble_id, key) + + if df.empty: + return df + + return df.astype( + { + "batch_id": int, + "control_name": str, + key: float, + } + ) - @staticmethod - @lru_cache(maxsize=128) def data_for_controls( - ensemble_id: str, parameter_keys: tuple[str, ...], ens_path: Path + self, ensemble_id: str, parameter_keys: tuple[str, ...] + ) -> pd.DataFrame: + return self._cached_controls(ensemble_id, parameter_keys) + + def _fetch_controls( + self, ensemble_id: str, parameter_keys: tuple[str, ...] ) -> pd.DataFrame: frames = [] for parameter_key in parameter_keys: - df = PlotApi.data_for_parameter(ensemble_id, parameter_key, ens_path) + df = self.data_for_parameter(ensemble_id, parameter_key) if not df.empty and {"batch_id", "realization"}.issubset(df.columns): value_cols = [ c for c in df.columns if c not in {"batch_id", "realization"} @@ -394,78 +338,58 @@ def data_for_controls( return pd.DataFrame() return pd.concat(frames, ignore_index=True) - @staticmethod - @lru_cache(maxsize=256) - def data_for_parameter( - ensemble_id: str, parameter_key: str, ens_path: Path - ) -> pd.DataFrame: - with create_ertserver_client(ens_path) as client: - http_response = client.get( - f"/ensembles/{ensemble_id}/parameters/{PlotApi.escape(parameter_key)}", - headers={"accept": "application/x-parquet"}, - timeout=TIMEOUT, - ) - PlotApi._check_http_response(http_response) + def data_for_parameter(self, ensemble_id: str, parameter_key: str) -> pd.DataFrame: + return self._cached_parameter(ensemble_id, parameter_key) - stream = io.BytesIO(http_response.content) - df = pd.read_parquet(stream) + def _fetch_parameter(self, ensemble_id: str, parameter_key: str) -> pd.DataFrame: + df = self._client.parameter(ensemble_id, parameter_key) - if {"batch_id", "realization"}.issubset(df.columns): - return df + if {"batch_id", "realization"}.issubset(df.columns): + return df - try: - df.columns = pd.to_datetime(df.columns, format="%Y-%m-%d %H:%M:%S") - except (ParserError, ValueError): - df.columns = [int(s) for s in df.columns] + try: + df.columns = pd.to_datetime(df.columns, format="%Y-%m-%d %H:%M:%S") + except (ParserError, ValueError): + df.columns = [int(s) for s in df.columns] - for col in df.columns: - if is_numeric_dtype(df[col]): - df[col] = df[col].astype(float) - return df + for col in df.columns: + if is_numeric_dtype(df[col]): + df[col] = df[col].astype(float) + return df def observation_locations(self) -> pd.DataFrame: - with create_ertserver_client(self.ens_path) as client: - http_response = client.get("/experiments", timeout=TIMEOUT) - self._check_http_response(http_response) - experiments = http_response.json() - - all_observations_list = [] - for experiment in experiments: - experiment_id = str(experiment["id"]) - http_response = client.get( - f"/experiments/{experiment_id}/observations", - timeout=TIMEOUT, - ) - self._check_http_response(http_response) - observations = http_response.json() + all_observations_list = [] + for experiment in self._client.experiments(): + experiment_id = str(experiment["id"]) + observations = self._client.experiment_observations(experiment_id) - try: - if not observations: - continue - new_obs = pd.concat( - ( - pd.DataFrame( - { - "east": obs["east"], - "north": obs["north"], - "radius": obs["radius"], - } - ) - for obs in observations - ), - ignore_index=True, - ).dropna() - if not new_obs.empty: - all_observations_list.append(new_obs) - except KeyError as e: - raise httpx.RequestError( - f"Observation payload missing coordinate key {e} for" - f" experiment {experiment_id}" - ) from e - - if not all_observations_list: - return pd.DataFrame() - return pd.concat(all_observations_list, ignore_index=True) + try: + if not observations: + continue + new_obs = pd.concat( + ( + pd.DataFrame( + { + "east": obs["east"], + "north": obs["north"], + "radius": obs["radius"], + } + ) + for obs in observations + ), + ignore_index=True, + ).dropna() + if not new_obs.empty: + all_observations_list.append(new_obs) + except KeyError as e: + raise httpx.RequestError( + f"Observation payload missing coordinate key {e} for" + f" experiment {experiment_id}" + ) from e + + if not all_observations_list: + return pd.DataFrame() + return pd.concat(all_observations_list, ignore_index=True) def observations_for_key(self, ensemble_ids: list[str], key: str) -> pd.DataFrame: """Returns a pandas DataFrame with the datapoints for a given observation key @@ -490,53 +414,45 @@ def observations_for_key(self, ensemble_ids: list[str], key: str) -> pd.DataFram if "@" in actual_response_key: actual_response_key = key.split("@", maxsplit=1)[0] filter_on = key_def.filter_on - with create_ertserver_client(self.ens_path) as client: - http_response = client.get( - f"/ensembles/{ensemble.id}/responses/{PlotApi.escape(actual_response_key)}/observations", - timeout=TIMEOUT, - params={"filter_on": json.dumps(filter_on)} - if filter_on is not None - else None, - ) - self._check_http_response(http_response) - - observations = http_response.json() + observations = self._client.response_observations( + ensemble.id, actual_response_key, filter_on + ) + try: + observations_dfs = [] + if not observations: + continue + + observations[0] # Just preserving the old logic/behavior + # but this should really be revised + except (KeyError, IndexError) as e: + raise httpx.RequestError( + f"Observation schema might have changed key={key}, " + f"ensemble_name={ensemble.name}, e={e}" + ) from e + + key_index: list[int | float | pd.Timestamp] + for obs in observations: try: - observations_dfs = [] - if not observations: - continue - - observations[0] # Just preserving the old logic/behavior - # but this should really be revised - except (KeyError, IndexError) as e: - raise httpx.RequestError( - f"Observation schema might have changed key={key}, " - f"ensemble_name={ensemble.name}, e={e}" - ) from e - - key_index: list[int | float | pd.Timestamp] - for obs in observations: + int(obs["x_axis"][0]) + key_index = [int(v) for v in obs["x_axis"]] + except ValueError: try: - int(obs["x_axis"][0]) - key_index = [int(v) for v in obs["x_axis"]] + float(obs["x_axis"][0]) + key_index = [float(v) for v in obs["x_axis"]] except ValueError: - try: - float(obs["x_axis"][0]) - key_index = [float(v) for v in obs["x_axis"]] - except ValueError: - key_index = [pd.Timestamp(v) for v in obs["x_axis"]] - - observations_dfs.append( - pd.DataFrame( - { - "STD": obs["errors"], - "OBS": obs["values"], - "key_index": key_index, - } - ) + key_index = [pd.Timestamp(v) for v in obs["x_axis"]] + + observations_dfs.append( + pd.DataFrame( + { + "STD": obs["errors"], + "OBS": obs["values"], + "key_index": key_index, + } ) + ) - all_observations = pd.concat([all_observations, *observations_dfs]) + all_observations = pd.concat([all_observations, *observations_dfs]) return all_observations.T @@ -580,14 +496,4 @@ def std_dev_for_parameter( if not ensemble: return np.array([]) - with create_ertserver_client(self.ens_path) as client: - http_response = client.get( - f"/ensembles/{ensemble.id}/parameters/{PlotApi.escape(key)}/std_dev", - params={"z": z}, - timeout=TIMEOUT, - ) - - if http_response.status_code == 200: - # Deserialize the numpy array - return np.load(io.BytesIO(http_response.content)) - return np.array([]) + return self._client.parameter_std_dev(ensemble.id, key, z) diff --git a/src/ert/gui/plotting/plot_window.py b/src/ert/gui/plotting/plot_window.py index 7b6dc18e549..ce78d88e8f0 100644 --- a/src/ert/gui/plotting/plot_window.py +++ b/src/ert/gui/plotting/plot_window.py @@ -167,7 +167,6 @@ def __init__( self.setWindowTitle(f"Plotting - {config_file}") self.activateWindow() self._preferred_ensemble_x_axis_format = PlotContext.INDEX_AXIS - self._ens_path = ens_path self._api = PlotApi(ens_path) self.local_version = get_storage_api_version() @@ -481,9 +480,7 @@ def fetch_data( try: # ruff: ignore[too-many-statements-in-try-clause] data = None if is_gradient_plot: - data = PlotApi.data_for_gradient( - ensemble.id, key, self._ens_path - ) + data = self._api.data_for_gradient(ensemble.id, key) elif ( key_def.response is not None or key_def.metadata.get("data_origin") @@ -495,20 +492,18 @@ def fetch_data( filter_on=key_def.filter_on, ) elif is_controls_plot: - data = PlotApi.data_for_controls( + data = self._api.data_for_controls( ensemble_id=ensemble.id, parameter_keys=tuple(selected_controls) or tuple(self._everest_parameters), - ens_path=self._ens_path, ) elif key_def.parameter is not None and ( key_def.parameter.type in {"gen_kw", "everest_parameters", "everest_objective"} ): - data = PlotApi.data_for_parameter( + data = self._api.data_for_parameter( ensemble_id=ensemble.id, parameter_key=key_def.parameter.name, - ens_path=self._ens_path, ) except BaseException as e: return ensemble, e diff --git a/src/ert/services/ert_client.py b/src/ert/services/ert_client.py index f8188fa2d15..0388331461d 100644 --- a/src/ert/services/ert_client.py +++ b/src/ert/services/ert_client.py @@ -14,7 +14,7 @@ import httpx import numpy as np import numpy.typing as npt -import polars as pl +import pandas as pd from .shared_client import Methods, SharedClient @@ -31,9 +31,8 @@ def _escape(value: str) -> str: def _uncached_copy[T](value: T) -> T: - """Polars frames are cheap to clone, and callers must not reach the cached one.""" - if isinstance(value, pl.DataFrame): - return cast("T", value.clone()) + if isinstance(value, pd.DataFrame): + return cast("T", value.copy()) return deepcopy(value) @@ -81,8 +80,8 @@ def _checked(response: httpx.Response) -> httpx.Response: return response -def _response_to_parquet(response: httpx.Response) -> pl.DataFrame: - return pl.read_parquet(io.BytesIO(response.content)) +def _response_to_parquet(response: httpx.Response) -> pd.DataFrame: + return pd.read_parquet(io.BytesIO(response.content)) class ErtClient: @@ -149,11 +148,11 @@ def ensemble_blobs(self, ensemble_id: str) -> list[dict[str, Any]]: def ensemble_blob(self, ensemble_id: str, uri: str) -> bytes: return self._get(f"/ensembles/{ensemble_id}/blobs/{_escape(uri)}").content - def parameter(self, ensemble_id: str, parameter_key: str) -> pl.DataFrame: + def parameter(self, ensemble_id: str, parameter_key: str) -> pd.DataFrame: return self._parameter(ensemble_id, parameter_key) @_cached - def _parameter(self, ensemble_id: str, parameter_key: str) -> pl.DataFrame: + def _parameter(self, ensemble_id: str, parameter_key: str) -> pd.DataFrame: return _response_to_parquet( self._get( f"/ensembles/{ensemble_id}/parameters/{_escape(parameter_key)}", @@ -178,7 +177,7 @@ def ert_response( ensemble_id: str, response_key: str, filter_on: dict[str, Any] | None = None, - ) -> pl.DataFrame: + ) -> pd.DataFrame: return _response_to_parquet( self._get( f"/ensembles/{ensemble_id}/responses/{_escape(response_key)}", @@ -187,11 +186,11 @@ def ert_response( ) ) - def gradient(self, ensemble_id: str, response_key: str) -> pl.DataFrame: + def gradient(self, ensemble_id: str, response_key: str) -> pd.DataFrame: return self._gradient(ensemble_id, response_key) @_cached - def _gradient(self, ensemble_id: str, response_key: str) -> pl.DataFrame: + def _gradient(self, ensemble_id: str, response_key: str) -> pd.DataFrame: return _response_to_parquet( self._request( "GET", diff --git a/src/ert/services/shared_client.py b/src/ert/services/shared_client.py index bbada0b8c12..c541c669f51 100644 --- a/src/ert/services/shared_client.py +++ b/src/ert/services/shared_client.py @@ -59,6 +59,13 @@ def get_client( cls._instance = cls(key, client) return cls._instance + @classmethod + def close_client(cls) -> None: + with cls._instance_lock: + if cls._instance is not None: + cls._instance._client.close() + cls._instance = None + @property def project(self) -> Path: return self._project diff --git a/tests/ert/performance_tests/test_dark_storage_performance.py b/tests/ert/performance_tests/test_dark_storage_performance.py index c1712021d90..93c9a2de3ce 100644 --- a/tests/ert/performance_tests/test_dark_storage_performance.py +++ b/tests/ert/performance_tests/test_dark_storage_performance.py @@ -24,8 +24,9 @@ from ert.dark_storage.endpoints import ensembles, experiments from ert.dark_storage.endpoints.observations import get_observations_for_response from ert.dark_storage.endpoints.responses import get_response -from ert.gui.plotting import plot_api from ert.gui.plotting.plot_api import PlotApi +from ert.services import ert_client +from ert.services.ert_client import ErtClient from ert.storage import Storage, open_storage @@ -54,7 +55,17 @@ def get_response_autofilter( @pytest.fixture(autouse=True) def use_testclient(monkeypatch): client = TestClient(app) - monkeypatch.setattr(plot_api, "create_ertserver_client", lambda project: client) + + class TestClientAdapter: + def request(self, method, url, **kwargs): + kwargs.pop("timeout", None) + return client.request(method, url, **kwargs) + + monkeypatch.setattr( + ErtClient, + "get_client", + classmethod(lambda cls, *args, **kwargs: cls(TestClientAdapter())), + ) def test_escape(s: str) -> str: """ @@ -63,7 +74,7 @@ def test_escape(s: str) -> str: """ return quote(quote(quote(s, safe=""))) - PlotApi.escape = test_escape + monkeypatch.setattr(ert_client, "_escape", test_escape) def run_in_loop[T](coro: Awaitable[T]) -> T: @@ -370,10 +381,9 @@ def run(): # Cycle through all ensembles and get all responses for key_info in key_infos_params: for ensemble in all_ensembles: - PlotApi.data_for_parameter( + api.data_for_parameter( ensemble_id=ensemble.id, parameter_key=key_info.parameter.name, - ens_path=api.ens_path, ) for key_info in key_infos_responses: diff --git a/tests/ert/ui_tests/gui/conftest.py b/tests/ert/ui_tests/gui/conftest.py index 1c301d4ecfc..4ca78e3c720 100644 --- a/tests/ert/ui_tests/gui/conftest.py +++ b/tests/ert/ui_tests/gui/conftest.py @@ -33,6 +33,7 @@ from ert.gui.tools.manage_experiments.storage_widget import AddWidget, StorageWidget from ert.plugins import get_site_plugins from ert.run_models import EnsembleExperiment, MultipleDataAssimilation +from ert.services import SharedClient from ert.storage import Storage from tests.ert.handle_run_path_dialog import handle_run_path_dialog @@ -47,6 +48,13 @@ def setup_svg_search_path(): ) +@pytest.fixture(autouse=True) +def reset_ert_api_client(): + # The client is process-wide and bound to one project, but each test has its own. + yield + SharedClient.close_client() + + @contextmanager def open_gui_with_config(config_path) -> Iterator[ErtMainWindow]: with ( diff --git a/tests/ert/unit_tests/gui/tools/plot/conftest.py b/tests/ert/unit_tests/gui/tools/plot/conftest.py index d549e8551a4..969018f3ff4 100644 --- a/tests/ert/unit_tests/gui/tools/plot/conftest.py +++ b/tests/ert/unit_tests/gui/tools/plot/conftest.py @@ -1,14 +1,12 @@ import io import os import shutil -from contextlib import contextmanager -from unittest.mock import MagicMock import pandas as pd import pytest -from ert.gui.plotting import plot_api from ert.gui.plotting.plot_api import PlotApi +from ert.services.ert_client import ErtClient class MockResponse: @@ -29,20 +27,19 @@ def is_success(self): return self.status_code == 200 -@pytest.fixture -def api(tmpdir, source_root, monkeypatch): - @contextmanager - def session(project: str): - yield MagicMock(get=mocked_requests_get) +class MockClient: + def request(self, method, url, **kwargs): + return mocked_requests_get(url, **kwargs) - monkeypatch.setattr(plot_api, "create_ertserver_client", session) +@pytest.fixture +def api(tmpdir, source_root): with tmpdir.as_cwd(): test_data_root = source_root / "test-data" / "ert" test_data_dir = test_data_root / "snake_oil" shutil.copytree(test_data_dir, "test_data") os.chdir("test_data") # ruff: ignore[banned-api] inside tmpdir.as_cwd() which restores - yield PlotApi(test_data_dir) + yield PlotApi(test_data_dir, ErtClient(MockClient())) def mocked_requests_get(*args, **kwargs): diff --git a/tests/ert/unit_tests/gui/tools/plot/test_plot_api.py b/tests/ert/unit_tests/gui/tools/plot/test_plot_api.py index 6056fb325db..bdbab2f5756 100644 --- a/tests/ert/unit_tests/gui/tools/plot/test_plot_api.py +++ b/tests/ert/unit_tests/gui/tools/plot/test_plot_api.py @@ -16,8 +16,8 @@ from ert.config import EverestObjectivesConfig, GenKwConfig, SummaryConfig from ert.dark_storage import common from ert.dark_storage.app import app -from ert.gui.plotting import plot_api from ert.gui.plotting.plot_api import PlotApi, PlotApiKeyDefinition +from ert.services import ErtClient, ert_client from ert.storage import open_storage from tests.ert.unit_tests.gui.tools.plot.conftest import MockResponse @@ -25,7 +25,17 @@ @pytest.fixture(autouse=True) def use_testclient(monkeypatch): client = TestClient(app) - monkeypatch.setattr(plot_api, "create_ertserver_client", lambda project: client) + + class TestClientAdapter: + def request(self, method, url, **kwargs): + kwargs.pop("timeout", None) + return client.request(method, url, **kwargs) + + monkeypatch.setattr( + ErtClient, + "get_client", + classmethod(lambda cls, *args, **kwargs: cls(TestClientAdapter())), + ) def test_escape(s: str) -> str: """ @@ -34,15 +44,7 @@ def test_escape(s: str) -> str: """ return quote(quote(quote(s, safe=""))) - PlotApi.escape = test_escape - - original_get = client.get - - def get_without_timeout(*args, **kwargs): - kwargs.pop("timeout", None) - return original_get(*args, **kwargs) - - client.get = get_without_timeout + monkeypatch.setattr(ert_client, "_escape", test_escape) def test_key_def_structure(api: PlotApi): @@ -123,9 +125,7 @@ def test_case_structure(api): def test_can_load_data_and_observations(api): responses = [("BPR:1,3,8", False), ("FOPR", True)] ensemble = next(x for x in api.get_all_ensembles() if x.name == "default_0") - data = PlotApi.data_for_parameter( - ensemble.id, "SNAKE_OIL_PARAM:BPR_138_PERSISTENCE", api.ens_path - ) + data = api.data_for_parameter(ensemble.id, "SNAKE_OIL_PARAM:BPR_138_PERSISTENCE") assert not data.empty for key, has_observations in responses: @@ -276,7 +276,7 @@ def test_plot_api_handles_urlescape(api_and_storage): def test_plot_api_handles_empty_gen_kw(api_and_storage): - _, storage, ens_path = api_and_storage + api, storage, _ = api_and_storage key = "gen_kw" name = "" experiment = storage.create_experiment( @@ -295,7 +295,7 @@ def test_plot_api_handles_empty_gen_kw(api_and_storage): } ) ensemble = storage.create_ensemble(experiment.id, ensemble_size=10) - assert PlotApi.data_for_parameter(str(ensemble.id), key, ens_path).empty + assert api.data_for_parameter(str(ensemble.id), key).empty ensemble.save_parameters( dataset=pl.DataFrame( { @@ -304,7 +304,7 @@ def test_plot_api_handles_empty_gen_kw(api_and_storage): } ), ) - assert PlotApi.data_for_parameter(str(ensemble.id), name, ens_path).to_csv() == ( + assert api.data_for_parameter(str(ensemble.id), name).to_csv() == ( dedent( """\ Realization,0 @@ -315,7 +315,7 @@ def test_plot_api_handles_empty_gen_kw(api_and_storage): def test_plot_api_handles_non_existant_gen_kw(api_and_storage): - _, storage, ens_path = api_and_storage + api, storage, _ = api_and_storage experiment = storage.create_experiment( experiment_config={ "parameter_configuration": [ @@ -331,14 +331,12 @@ def test_plot_api_handles_non_existant_gen_kw(api_and_storage): } ) ensemble = storage.create_ensemble(experiment.id, ensemble_size=10) - assert PlotApi.data_for_parameter(str(ensemble.id), "gen_kw", ens_path).empty - assert PlotApi.data_for_parameter( - str(ensemble.id), "gen_kw:does_not_exist", ens_path - ).empty + assert api.data_for_parameter(str(ensemble.id), "gen_kw").empty + assert api.data_for_parameter(str(ensemble.id), "gen_kw:does_not_exist").empty def test_plot_api_handles_colons_in_parameter_keys(api_and_storage): - _, storage, ens_path = api_and_storage + api, storage, _ = api_and_storage experiment = storage.create_experiment( experiment_config={ "parameter_configuration": [ @@ -363,7 +361,7 @@ def test_plot_api_handles_colons_in_parameter_keys(api_and_storage): } ), ) - test = PlotApi.data_for_parameter(str(ensemble.id), "subgroup:1:2:2", ens_path) + test = api.data_for_parameter(str(ensemble.id), "subgroup:1:2:2") assert test.to_numpy() == np.array([[10]]) @@ -676,51 +674,41 @@ def test_that_data_for_response_returns_empty_for_gradient_only_ensemble( def test_that_data_for_gradient_returns_empty_when_no_gradient_parquet_saved( api_and_storage, ): - _, storage, ens_path = api_and_storage + api, storage, _ = api_and_storage ensemble, objective_key = _create_gradient_only_ensemble(storage) # No batch_objective_gradient.parquet was saved, so the gradient # endpoint returns an empty DataFrame. PlotApi must not crash. - result = PlotApi.data_for_gradient(str(ensemble.id), objective_key, ens_path) + result = api.data_for_gradient(str(ensemble.id), objective_key) assert isinstance(result, pd.DataFrame) assert result.empty def test_that_data_for_gradient_is_fetched_once_for_repeated_calls(api_and_storage): - _, storage, ens_path = api_and_storage + api, storage, _ = api_and_storage ensemble, objective_key = _create_gradient_only_ensemble(storage) - PlotApi.data_for_gradient.cache_clear() # type: ignore - for i in range(5): # The addition of @ should result in the same cache key # if the prefix is the same. key = objective_key + "@filler" if i % 2 == 0 else objective_key - _ = PlotApi.data_for_gradient(str(ensemble.id), key, ens_path) + _ = api.data_for_gradient(str(ensemble.id), key) # hits, misses, maxsize, currsize expected_cache_info = (4, 1, 256, 1) - assert ( - PlotApi.data_for_gradient.cache_info() == expected_cache_info # type: ignore - ) + assert api._cached_gradient.cache_info() == expected_cache_info def test_that_data_for_controls_and_parameters_is_fetched_once_for_repeated_calls(api): - - PlotApi.data_for_controls.cache_clear() - PlotApi.data_for_parameter.cache_clear() - ensemble = next(x for x in api.get_all_ensembles() if x.name == "default_0") for _ in range(5): - _ = PlotApi.data_for_controls( - ensemble.id, ("SNAKE_OIL_PARAM:BPR_138_PERSISTENCE",), api.ens_path - ) + _ = api.data_for_controls(ensemble.id, ("SNAKE_OIL_PARAM:BPR_138_PERSISTENCE",)) # hits, misses, maxsize, currsize expected_cache_info_controls = (4, 1, 128, 1) expected_cache_info_parameters = (0, 1, 256, 1) - assert PlotApi.data_for_controls.cache_info() == expected_cache_info_controls - assert PlotApi.data_for_parameter.cache_info() == expected_cache_info_parameters + assert api._cached_controls.cache_info() == expected_cache_info_controls + assert api._cached_parameter.cache_info() == expected_cache_info_parameters diff --git a/tests/ert/unit_tests/gui/tools/plot/test_plot_window.py b/tests/ert/unit_tests/gui/tools/plot/test_plot_window.py index 61a4b16213c..28c7a7a0c5b 100644 --- a/tests/ert/unit_tests/gui/tools/plot/test_plot_window.py +++ b/tests/ert/unit_tests/gui/tools/plot/test_plot_window.py @@ -255,14 +255,12 @@ def test_that_plotting_gen_kw_parameter_with_negative_values_hides_log_scale_che plot_api_key_def_negative, ] - def mock_data_for_parameter( - ensemble_id: str, parameter_key: str, ens_path: Path - ) -> pd.DataFrame: + def mock_data_for_parameter(ensemble_id: str, parameter_key: str) -> pd.DataFrame: if parameter_key == "gen_kw_a": return pd.DataFrame({0: [0.1, 0.5, 0.9]}) return pd.DataFrame({0: [-0.1, -0.5, -0.9]}) - mock_plot_api_cls.data_for_parameter.side_effect = mock_data_for_parameter + mock_plot_api.data_for_parameter.side_effect = mock_data_for_parameter mock_plot_api.has_history_data.return_value = False mock_plot_api.get_all_ensembles.return_value = [ @@ -643,9 +641,7 @@ def test_that_log_scale_state_is_preserved_when_switching_plot_tabs( mock_plot_api.responses_api_key_defs = [] mock_plot_api.parameters_api_key_defs = [key_def] - mock_plot_api_cls.data_for_parameter.return_value = pd.DataFrame( - {0: [0.1, 0.5, 0.9]} - ) + mock_plot_api.data_for_parameter.return_value = pd.DataFrame({0: [0.1, 0.5, 0.9]}) mock_plot_api.has_history_data.return_value = False mock_plot_api.get_all_ensembles.return_value = [ EnsembleObject( @@ -840,7 +836,7 @@ def test_that_density_tabs_show_log_scale_only_for_valid_gen_kw_values( mock_plot_api.responses_api_key_defs = [] mock_plot_api.parameters_api_key_defs = [key_def] - mock_plot_api_cls.data_for_parameter.return_value = pd.DataFrame({0: values}) + mock_plot_api.data_for_parameter.return_value = pd.DataFrame({0: values}) mock_plot_api.has_history_data.return_value = False mock_plot_api.get_all_ensembles.return_value = [ EnsembleObject( @@ -912,7 +908,7 @@ def test_that_plot_window_ignores_negative_check_for_non_numeric_columns( mock_plot_api.parameters_api_key_defs = [plot_api_key_def] def mixed_dtype_data_for_parameter( - ensemble_id: str, parameter_key: str, ens_path: Path + ensemble_id: str, parameter_key: str ) -> pd.DataFrame: assert parameter_key == "animal_type" return pd.DataFrame( @@ -921,7 +917,7 @@ def mixed_dtype_data_for_parameter( } ) - mock_plot_api_cls.data_for_parameter.side_effect = mixed_dtype_data_for_parameter + mock_plot_api.data_for_parameter.side_effect = mixed_dtype_data_for_parameter mock_plot_api.has_history_data.return_value = False mock_plot_api.get_all_ensembles.return_value = [ EnsembleObject( @@ -1320,9 +1316,7 @@ def test_that_clearing_custom_title_restores_key_title_when_rendering( "2026-01-01T00:00:00", ) ] - mock_plot_api_cls.data_for_parameter.return_value = pd.DataFrame( - {0: [1.0, 2.0, 3.0]} - ) + mock_plot_api.data_for_parameter.return_value = pd.DataFrame({0: [1.0, 2.0, 3.0]}) mock_plot_api.has_history_data.return_value = False monkeypatch.setattr( diff --git a/tests/ert/unit_tests/services/test_ert_client.py b/tests/ert/unit_tests/services/test_ert_client.py index 61994085093..3f7c392f324 100644 --- a/tests/ert/unit_tests/services/test_ert_client.py +++ b/tests/ert/unit_tests/services/test_ert_client.py @@ -1,7 +1,7 @@ import io from typing import Any -import polars as pl +import pandas as pd import pytest from ert.services.ert_client import ErtClient @@ -41,7 +41,7 @@ def count_requests(self, fragment: str) -> int: def _payload_for(url: str) -> Any: if "/parameters/" in url or "/responses/" in url or "/gradients/" in url: stream = io.BytesIO() - pl.DataFrame({"0": [1.0, 2.0, 3.0]}).write_parquet(stream) + pd.DataFrame({"0": [1.0, 2.0, 3.0]}).to_parquet(stream) return stream.getvalue() if url == "/experiments": return [{"id": "exp_1", "ensemble_ids": ["ens_1"]}]