Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
640 changes: 273 additions & 367 deletions src/ert/gui/plotting/plot_api.py

Large diffs are not rendered by default.

11 changes: 3 additions & 8 deletions src/ert/gui/plotting/plot_window.py
Original file line number Diff line number Diff line change
Expand Up @@ -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()
Expand Down Expand Up @@ -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")
Expand All @@ -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
Expand Down
21 changes: 10 additions & 11 deletions src/ert/services/ert_client.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand All @@ -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())
Comment thread
eqbech marked this conversation as resolved.
return deepcopy(value)


Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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:
Comment thread
eqbech marked this conversation as resolved.
return _response_to_parquet(
self._get(
f"/ensembles/{ensemble_id}/parameters/{_escape(parameter_key)}",
Expand All @@ -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)}",
Expand All @@ -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",
Expand Down
7 changes: 7 additions & 0 deletions src/ert/services/shared_client.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
20 changes: 15 additions & 5 deletions tests/ert/performance_tests/test_dark_storage_performance.py
Original file line number Diff line number Diff line change
Expand Up @@ -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


Expand Down Expand Up @@ -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:
"""
Expand All @@ -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:
Expand Down Expand Up @@ -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:
Expand Down
8 changes: 8 additions & 0 deletions tests/ert/ui_tests/gui/conftest.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand All @@ -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 (
Expand Down
17 changes: 7 additions & 10 deletions tests/ert/unit_tests/gui/tools/plot/conftest.py
Original file line number Diff line number Diff line change
@@ -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:
Expand All @@ -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):
Expand Down
Loading