From d103a8a688cfd6d20f1bedfb191e6d6c037a11b3 Mon Sep 17 00:00:00 2001 From: Benjamin A Chandler Date: Tue, 8 Sep 2026 13:50:56 +0200 Subject: [PATCH] Replace-experiment-client-with-ert-client --- src/ert/dark_storage/common.py | 12 + .../endpoints/experiment_server.py | 2 +- src/ert/gui/experiments/experiment_client.py | 161 ----------- src/ert/services/ert_client.py | 208 ++++++++++++-- src/everest/bin/everest_script.py | 33 +-- src/everest/bin/kill_script.py | 9 +- src/everest/bin/monitor_script.py | 15 +- src/everest/bin/utils.py | 28 +- src/everest/detached/__init__.py | 14 - src/everest/detached/client.py | 223 ++------------- src/everest/detached/everserver.py | 5 +- src/everest/gui/main_window.py | 61 ++-- src/everest/strings.py | 13 - .../unit_tests/services/test_ert_client.py | 55 ++++ .../entry_points/test_everest_entry.py | 109 ++++--- tests/everest/entry_points/test_everexport.py | 0 tests/everest/test_detached.py | 99 +++++-- tests/everest/test_everest_client.py | 266 ------------------ tests/everest/test_everest_output.py | 4 - tests/everest/test_everserver.py | 9 +- tests/everest/test_monitor.py | 169 ++++++----- 21 files changed, 574 insertions(+), 921 deletions(-) delete mode 100644 src/ert/gui/experiments/experiment_client.py delete mode 100644 tests/everest/entry_points/test_everexport.py delete mode 100644 tests/everest/test_everest_client.py diff --git a/src/ert/dark_storage/common.py b/src/ert/dark_storage/common.py index 1459f731dc7..cb1bd3aaf4e 100644 --- a/src/ert/dark_storage/common.py +++ b/src/ert/dark_storage/common.py @@ -4,6 +4,7 @@ import re from collections.abc import Iterator from contextlib import contextmanager +from enum import StrEnum, auto from importlib import metadata import pandas as pd @@ -25,6 +26,17 @@ _storage: Storage | None = None +class EverEndpoints(StrEnum): + STOP = auto() + START_EXPERIMENT = auto() + CONFIG_PATH = auto() + START_TIME = auto() + EXPERIMENTS = auto() + STATUS = auto() + EVENTS = auto() + RUNPATH = auto() + + def get_storage() -> Storage: global _storage if _storage is None: diff --git a/src/ert/dark_storage/endpoints/experiment_server.py b/src/ert/dark_storage/endpoints/experiment_server.py index 7e01ef4028d..edfa488ed35 100644 --- a/src/ert/dark_storage/endpoints/experiment_server.py +++ b/src/ert/dark_storage/endpoints/experiment_server.py @@ -32,6 +32,7 @@ from ert.base_model_context import use_runtime_plugins from ert.config import ConfigWarning, QueueSystem +from ert.dark_storage.common import EverEndpoints from ert.ensemble_evaluator import EndEvent, EvaluatorServerConfig from ert.ensemble_evaluator.event import FullSnapshotEvent, SnapshotUpdateEvent from ert.ensemble_evaluator.snapshot import EnsembleSnapshot @@ -46,7 +47,6 @@ from everest.strings import ( OPT_FAILURE_ALL_REALIZATIONS, OPT_FAILURE_REALIZATIONS, - EverEndpoints, ) router = APIRouter(prefix="/experiment_server", tags=["experiment_server"]) diff --git a/src/ert/gui/experiments/experiment_client.py b/src/ert/gui/experiments/experiment_client.py deleted file mode 100644 index 59d38d08272..00000000000 --- a/src/ert/gui/experiments/experiment_client.py +++ /dev/null @@ -1,161 +0,0 @@ -from __future__ import annotations - -import logging -import queue -import ssl -import time -import traceback -from base64 import b64encode -from http import HTTPStatus -from pathlib import Path - -import requests -from pydantic import ValidationError -from requests import HTTPError -from websockets.exceptions import ConnectionClosedError -from websockets.sync.client import connect - -from _ert.threading import ErtThread -from ert.ensemble_evaluator import EvaluatorServerConfig -from ert.run_models import RunModelAPI -from ert.run_models.event import StatusEvents, status_event_from_json -from everest.strings import EverEndpoints - -logger = logging.getLogger(__name__) - - -class ExperimentClient: - def __init__( - self, - experiment_id: str, - url: str, - cert_file: str, - username: str, - password: str, - ssl_context: ssl.SSLContext, - ) -> None: - self._experiment_id = experiment_id - self._url = url - self._cert = cert_file - self._username = username - self._password = password - self._ssl_context = ssl_context - - self._is_alive = False - self._start_time: int | None = None - - def _http_get(self, endpoint: str) -> requests.Response: - return requests.get( - f"{self._url}/{endpoint}", - verify=self._cert, - auth=(self._username, self._password), - proxies={"http": None, "https": None}, # type: ignore - ) - - def _http_post(self, endpoint: str) -> requests.Response: - return requests.post( - f"{self._url}/{endpoint}", - verify=self._cert, - auth=(self._username, self._password), - proxies={"http": None, "https": None}, # type: ignore - ) - - @property - def config(self) -> dict[str, str]: - return self._http_get( - f"{EverEndpoints.CONFIG_PATH}/{self._experiment_id}" - ).json() - - @property - def credentials(self) -> str: - return b64encode(f"{self._username}:{self._password}".encode()).decode() - - def setup_event_queue_from_ws_endpoint( - self, - refresh_interval: float = 0.01, - open_timeout: float = 30, - websocket_recv_timeout: float = 1.0, - ) -> tuple[queue.SimpleQueue[StatusEvents], ErtThread]: - event_queue: queue.SimpleQueue[StatusEvents] = queue.SimpleQueue() - - def passthrough_ws_events() -> None: - try: # ruff: ignore[too-many-statements-in-try-clause] - with connect( - self._url.replace("https://", "wss://") - + f"/{EverEndpoints.EVENTS}/{self._experiment_id}", - ssl=self._ssl_context, - open_timeout=open_timeout, - additional_headers={"Authorization": f"Basic {self.credentials}"}, - ) as websocket: - while not self._is_alive: - try: - message = websocket.recv(timeout=websocket_recv_timeout) - except TimeoutError: - message = None - if message: - try: - event = status_event_from_json(message) - event_queue.put(event) - except ValidationError as e: - logger.error( - "Error when processing event %s", exc_info=e - ) - - time.sleep(refresh_interval) - except ConnectionClosedError: - logger.debug("Connection closed by server") - except Exception: - logger.debug(traceback.format_exc()) - - monitor_thread = ErtThread( - name="everest_gui_event_monitor", - target=passthrough_ws_events, - daemon=True, - ) - - return event_queue, monitor_thread - - def create_run_model_api(self) -> RunModelAPI: - def start_fn( - evaluator_server_config: EvaluatorServerConfig, - *, - rerun_failed_realizations: bool = False, - ) -> None: - pass - - return RunModelAPI( - experiment_name=Path(self.config["config_path"]).name, - supports_rerunning_failed_realizations=False, - start_simulations_thread=start_fn, - cancel=self.stop, - has_failed_realizations=lambda: False, - ) - - def stop(self) -> None: - try: - response = self._http_post(EverEndpoints.STOP) - except requests.exceptions.ConnectionError as e: - logger.error( - "Connection error when cancelling EVEREST " - f"experiment: {''.join(traceback.format_exception(e))}" - ) - print("Failed to cancel experiment") - return - except HTTPError as e: - logger.error( - "HTTP error when cancelling EVEREST " - f"experiment: {''.join(traceback.format_exception(e))}" - ) - print("Failed to cancel experiment") - return - if response.status_code == 200: - logger.info("Cancelled experiment from EVEREST") - print("Successfully cancelled experiment") - else: - logger.error( - f"Failed to cancel EVEREST experiment: " - f"POST @ {self._url}/{EverEndpoints.STOP}, " - f"server responded with status {response.status_code}: " - f"{HTTPStatus(response.status_code).phrase}" - ) - print("Failed to cancel experiment") diff --git a/src/ert/services/ert_client.py b/src/ert/services/ert_client.py index 8d738f9259c..f727e209ccb 100644 --- a/src/ert/services/ert_client.py +++ b/src/ert/services/ert_client.py @@ -2,28 +2,48 @@ import io import json +import logging +import queue +import ssl import threading +import time +import traceback +from base64 import b64encode from collections import OrderedDict -from collections.abc import Callable +from collections.abc import Callable, Generator from copy import deepcopy from functools import wraps from os import PathLike -from typing import Any, Concatenate, cast +from typing import TYPE_CHECKING, Any, Concatenate, cast from urllib.parse import quote import httpx import numpy as np import numpy.typing as npt import pandas as pd +from pydantic import ValidationError +from websockets.exceptions import ConnectionClosedError, ConnectionClosedOK +from websockets.sync.client import connect + +from _ert.threading import ErtThread +from ert.dark_storage.common import EverEndpoints from .shared_client import ErtClientConnectionInfo, Methods, SharedClient +if TYPE_CHECKING: + from ert.run_models.event import StatusEvents + DEFAULT_TIMEOUT = 120 DEFAULT_CACHE_SIZE = 256 +# Specifies how many times to try a http request within the specified timeout. +_HTTP_REQUEST_RETRY = 10 + _PARQUET = {"accept": "application/x-parquet"} _EXPERIMENT_SERVER = "/experiment_server" +logger = logging.getLogger(__name__) + def _escape(value: str) -> str: """Keys may contain slashes, and the server decodes the path segment once.""" @@ -132,6 +152,58 @@ def clear_cache(self) -> None: with self._cache_lock: self._cache.clear() + # <-------------- General -------------------> + + def server_is_running(self, *, timeout: float | None = None) -> bool: + try: + response = self._request( + "GET", + f"{_EXPERIMENT_SERVER}/", + auth=self._auth, + timeout=timeout or self._timeout, + ) + except Exception: + return False + return response.status_code == httpx.codes.OK + + def wait_for_server(self, timeout: float) -> None: + """ + Polls server availability until timeout (measured in seconds). + + Raises an exception if no response within the timeout. + """ + wait_start_time: float = time.monotonic() + while time.monotonic() - wait_start_time <= timeout: + if self.server_is_running(timeout=1): + return + until_timeout = max(0, timeout - (time.monotonic() - wait_start_time)) + time.sleep(min(1, until_timeout)) + raise RuntimeError( + "Failed to get reply from server " + f"within {time.monotonic() - wait_start_time:g} seconds" + ) + + def wait_for_server_to_stop( + self, timeout: float, attempts: int = _HTTP_REQUEST_RETRY + ) -> None: + """ + Checks server has stopped `attempts` times. Waits + progressively longer between each check. + + Raise an exception when the timeout is reached. + """ + if self.server_is_running(timeout=1): + sleep_time_increment = float(timeout) / (2**attempts - 1) + for retry_count in range(attempts): + sleep_time = sleep_time_increment * (2**retry_count) + time.sleep(sleep_time) + if not self.server_is_running(timeout=1): + return + + # If number of retries reached and server still running - throw exception + if self.server_is_running(timeout=1): + raise Exception("Failed to stop server within configured timeout.") + # <-------------- Dark Storage --------------> def healthcheck(self) -> str: @@ -141,7 +213,7 @@ def version(self) -> str: return str(self._get("/version").json()) def experiments(self) -> list[dict[str, Any]]: - return list(self._get("/experiments").json()) + return self._get("/experiments").json() def ensemble(self, ensemble_id: str) -> dict[str, Any]: return dict(self._get(f"/ensembles/{ensemble_id}").json()) @@ -223,49 +295,147 @@ def response_observations( # <------------- Experiment Server -------------> - def experiment_server_is_running(self) -> bool: - try: - response = self._request("GET", f"{_EXPERIMENT_SERVER}/", auth=self._auth) - except httpx.TransportError: - return False - return response.status_code == httpx.codes.OK - def experiment_ids(self) -> list[str]: - response = self._experiment_server_get("experiments") + response = self._experiment_server_get(EverEndpoints.EXPERIMENTS) return list(response.json()["experiment_ids"]) def experiment_status(self, experiment_id: str) -> dict[str, Any]: - return dict(self._experiment_server_get(f"status/{experiment_id}").json()) + return dict( + self._experiment_server_get( + f"{EverEndpoints.STATUS}/{experiment_id}" + ).json() + ) - def experiment_config_path(self, experiment_id: str) -> dict[str, Any]: - return dict(self._experiment_server_get(f"config_path/{experiment_id}").json()) + def experiment_config(self, experiment_id: str) -> dict[str, str]: + return self._experiment_server_get( + f"{EverEndpoints.CONFIG_PATH}/{experiment_id}" + ).json() def experiment_start_time(self, experiment_id: str) -> int: - return int(self._experiment_server_get(f"start_time/{experiment_id}").text) + return int( + self._experiment_server_get( + f"{EverEndpoints.START_TIME}/{experiment_id}" + ).text + ) def start_experiment(self, config: dict[str, Any]) -> str: response = self._request( "POST", - f"{_EXPERIMENT_SERVER}/start_experiment", + f"{_EXPERIMENT_SERVER}/{EverEndpoints.START_EXPERIMENT}", auth=self._auth, json=config, ) return str(_checked(response).json()["experiment_id"]) - def stop_experiment_server(self) -> None: - _checked(self._request("POST", f"{_EXPERIMENT_SERVER}/stop", auth=self._auth)) + def stop_experiment_server(self) -> bool: + return ( + self._request( + "POST", + f"{_EXPERIMENT_SERVER}/{EverEndpoints.STOP}", + auth=self._auth, + ).status_code + == 200 + ) def runpath_exists(self, paths: list[str]) -> bool: response = self._request( "POST", - f"{_EXPERIMENT_SERVER}/runpath", + f"{_EXPERIMENT_SERVER}/{EverEndpoints.RUNPATH}", auth=self._auth, json={"paths": paths}, ) return response.status_code == httpx.codes.OK + # <-------------- WebSocket --------------> + + def iter_events( + self, + experiment_id: str, + refresh_interval: float = 0.01, + open_timeout: float = 30, + websocket_recv_timeout: float = 1.0, + ) -> Generator[StatusEvents, None, None]: + """Yield events synchronously until the WebSocket disconnects. + + Each iterator owns a separate connection and blocks only its consuming + thread. Close the iterator when stopping consumption early. + """ + from ert.run_models.event import ( # ruff: ignore[import-outside-top-level] + status_event_from_json, + ) + + url = ( + self.conn_info.base_url.replace("https://", "wss://") + + f"{_EXPERIMENT_SERVER}/{EverEndpoints.EVENTS}/{experiment_id}" + ) + username, password = self._auth + credentials = b64encode(f"{username}:{password}".encode()).decode() + + try: # ruff: ignore[too-many-statements-in-try-clause] + with connect( + url, + ssl=self._ssl_context, + open_timeout=open_timeout, + additional_headers={"Authorization": f"Basic {credentials}"}, + ) as websocket: + while True: + try: + message = websocket.recv(timeout=websocket_recv_timeout) + except TimeoutError: + message = None + if message: + try: + event = status_event_from_json(message) + except ValidationError as e: + logger.error("Error when processing event %s", exc_info=e) + else: + yield event + + time.sleep(refresh_interval) + except (ConnectionClosedOK, ConnectionClosedError): + logger.debug("Connection closed by server") + except Exception: + logger.error(traceback.format_exc()) + + def setup_event_queue_from_ws_endpoint( + self, + experiment_id: str, + refresh_interval: float = 0.01, + open_timeout: float = 30, + websocket_recv_timeout: float = 1.0, + ) -> tuple[queue.SimpleQueue[StatusEvents], ErtThread]: + """Return a queue of experiment events and the thread that fills it. + + The caller owns the thread and must start it. + """ + event_queue: queue.SimpleQueue[StatusEvents] = queue.SimpleQueue() + + def passthrough_ws_events() -> None: + for event in self.iter_events( + experiment_id, + refresh_interval=refresh_interval, + open_timeout=open_timeout, + websocket_recv_timeout=websocket_recv_timeout, + ): + event_queue.put(event) + + monitor_thread = ErtThread( + name="ert_storage_api_event_monitor", + target=passthrough_ws_events, + daemon=True, + ) + + return event_queue, monitor_thread + # <-------------- Internals --------------> + @property + def _ssl_context(self) -> ssl.SSLContext | None: + cert = self._client.conn_info.cert + if not isinstance(cert, str): + return None + return ssl.create_default_context(cafile=cert) + @property def _auth(self) -> tuple[str, str]: """Experiment-server routes authenticate with HTTP Basic, not the token diff --git a/src/everest/bin/everest_script.py b/src/everest/bin/everest_script.py index 688790e4b0c..f2556542518 100755 --- a/src/everest/bin/everest_script.py +++ b/src/everest/bin/everest_script.py @@ -22,9 +22,7 @@ from ert.utils import makedirs_if_needed from everest.config import EverestConfig, ServerConfig from everest.detached import ( - start_experiment, start_server, - wait_for_server, ) from everest.strings import EVEREST from everest.util import ( @@ -98,10 +96,16 @@ def everest_entry(args: list[str] | None = None) -> None: if threading.current_thread() is threading.main_thread(): signal.signal( signal.SIGINT, - partial(handle_keyboard_interrupt, options=options), + partial(signal.default_int_handler), ) - asyncio.run(run_everest(options)) + async def run_with_interrupt_handler() -> None: + try: + await run_everest(options) + except KeyboardInterrupt: + handle_keyboard_interrupt(signal.SIGINT, None, options) + + asyncio.run(run_with_interrupt_handler()) def _build_args_parser() -> argparse.ArgumentParser: @@ -231,18 +235,14 @@ async def directory_is_nonempty(path: Path) -> bool: client = ErtClient.get_client( Path(ServerConfig.get_session_dir(options.config.output_dir)) ) - wait_for_server(client, timeout=600) + client.wait_for_server(timeout=600) print("EVEREST server found!") logger.info( "Got response from everserver after " f"waiting for {time.monotonic() - wait_start_time:g} seconds. " "Starting experiment" ) - - experiment_id = start_experiment( - server_context=ServerConfig.get_server_context_from_conn_info(client.conn_info), - config=options.config, - ) + experiment_id = client.start_experiment(config_dict) # blocks until the run is finished if options.gui: @@ -253,10 +253,7 @@ async def directory_is_nonempty(path: Path) -> bool: if options.disable_monitoring else run_detached_monitor, name="EVEREST CLI monitor thread", - args=[ - ServerConfig.get_server_context_from_conn_info(client.conn_info), - experiment_id, - ], + args=[client, experiment_id], daemon=True, ) monitor_thread.start() @@ -264,16 +261,12 @@ async def directory_is_nonempty(path: Path) -> bool: monitor_thread.join() elif options.disable_monitoring: run_empty_detached_monitor( - server_context=ServerConfig.get_server_context_from_conn_info( - client.conn_info - ), + client=client, experiment_id=experiment_id, ) else: run_detached_monitor( - server_context=ServerConfig.get_server_context_from_conn_info( - client.conn_info - ), + client=client, experiment_id=experiment_id, ) diff --git a/src/everest/bin/kill_script.py b/src/everest/bin/kill_script.py index 1b0e43cc331..6079a463a64 100755 --- a/src/everest/bin/kill_script.py +++ b/src/everest/bin/kill_script.py @@ -15,7 +15,6 @@ from ert.services.ert_client import ErtClient from everest.bin.utils import setup_logging from everest.config import EverestConfig, ServerConfig -from everest.detached import stop_server, wait_for_server_to_stop from everest.util import version_info logger = logging.getLogger(__name__) @@ -78,15 +77,11 @@ def kill_everest(options: argparse.Namespace) -> None: Path(ServerConfig.get_session_dir(options.config.output_dir)), connect_timeout=1, ) - server_context = ServerConfig.get_server_context_from_conn_info( - client.conn_info - ) - except TimeoutError: print("Server is not running.") return - stopping = stop_server(server_context) + stopping = client.stop_experiment_server() if threading.current_thread() is threading.main_thread(): signal.signal(signal.SIGINT, partial(_handle_keyboard_interrupt, after=True)) @@ -95,7 +90,7 @@ def kill_everest(options: argparse.Namespace) -> None: return try: print("Waiting for server to stop ...") - wait_for_server_to_stop(server_context, timeout=60) + client.wait_for_server_to_stop(timeout=60) print("Server stopped.") except Exception: logger.debug(traceback.format_exc()) diff --git a/src/everest/bin/monitor_script.py b/src/everest/bin/monitor_script.py index 9f0d437e4a9..e7c59b9c2db 100755 --- a/src/everest/bin/monitor_script.py +++ b/src/everest/bin/monitor_script.py @@ -10,7 +10,6 @@ from ert.services.ert_client import ErtClient from ert.storage import ErtStorageException, ExperimentState from everest.config import EverestConfig, ServerConfig -from everest.detached.client import get_experiments from .utils import ( ArgParseFormatter, @@ -30,10 +29,13 @@ def monitor_entry(args: list[str] | None = None) -> None: if threading.current_thread() is threading.main_thread(): signal.signal( signal.SIGINT, - partial(handle_keyboard_interrupt, options=options), + partial(signal.default_int_handler), ) - monitor_everest(options) + try: + monitor_everest(options) + except KeyboardInterrupt: + handle_keyboard_interrupt(signal.SIGINT, None, options) def _build_args_parser() -> argparse.ArgumentParser: @@ -81,11 +83,8 @@ def monitor_everest(options: argparse.Namespace) -> None: client = ErtClient.get_client( Path(ServerConfig.get_session_dir(config.output_dir)), connect_timeout=1 ) - server_context = ServerConfig.get_server_context_from_conn_info( - client.conn_info - ) - experiment_id = get_experiments(server_context)[-1] - run_detached_monitor(server_context=server_context, experiment_id=experiment_id) + experiment_id = client.experiment_ids()[-1] + run_detached_monitor(client=client, experiment_id=experiment_id) try: experiment_status = get_experiment_status(str(config.storage_dir)) diff --git a/src/everest/bin/utils.py b/src/everest/bin/utils.py index e080aa69a5d..0e5f9959939 100644 --- a/src/everest/bin/utils.py +++ b/src/everest/bin/utils.py @@ -32,12 +32,7 @@ from ert.utils import makedirs_if_needed from everest.config import EverestConfig from everest.config.server_config import ServerConfig -from everest.detached import ( - server_is_running, - start_monitor, - stop_server, - wait_for_server_to_stop, -) +from everest.detached.client import start_monitor from everest.strings import EVEREST, OPT_PROGRESS_ID, SIM_PROGRESS_ID from everest.util import format_list @@ -107,14 +102,13 @@ def handle_keyboard_interrupt(signum: int, _: Any, options: argparse.Namespace) ) try: client = ErtClient.get_client( - Path(ServerConfig.get_session_dir(options.config.output_dir)) - ) - server_context = ServerConfig.get_server_context_from_conn_info( - client.conn_info + Path(ServerConfig.get_session_dir(options.config.output_dir)), + connect_timeout=1, ) - if server_is_running(*server_context): - stop_server(server_context) - wait_for_server_to_stop(server_context, timeout=10) + if client.server_is_running(timeout=1): + client.stop_experiment_server() + client.wait_for_server_to_stop(timeout=10) + print("Server stopped successfully.") except TimeoutError: print("No running server found.") @@ -396,18 +390,18 @@ def _clear(self) -> None: def run_detached_monitor( - server_context: tuple[str, str, tuple[str, str]], + client: ErtClient, experiment_id: str, ) -> None: monitor = _DetachedMonitor() - start_monitor(server_context, callback=monitor.update, experiment_id=experiment_id) + start_monitor(client, callback=monitor.update, experiment_id=experiment_id) def run_empty_detached_monitor( - server_context: tuple[str, str, tuple[str, str]], + client: ErtClient, experiment_id: str, ) -> None: - start_monitor(server_context, callback=lambda _: None, experiment_id=experiment_id) + start_monitor(client, callback=lambda _: None, experiment_id=experiment_id) def remove_show_scaling_warning_setting() -> None: diff --git a/src/everest/detached/__init__.py b/src/everest/detached/__init__.py index 1224682a9d7..9d4ec530a04 100644 --- a/src/everest/detached/__init__.py +++ b/src/everest/detached/__init__.py @@ -1,25 +1,11 @@ """Client methods for interacting with everserver""" from .client import ( - PROXY, - get_experiments, - server_is_running, - start_experiment, start_monitor, start_server, - stop_server, - wait_for_server, - wait_for_server_to_stop, ) __all__ = [ - "PROXY", - "get_experiments", - "server_is_running", - "start_experiment", "start_monitor", "start_server", - "stop_server", - "wait_for_server", - "wait_for_server_to_stop", ] diff --git a/src/everest/detached/client.py b/src/everest/detached/client.py index b62edc2ceae..ec7b17221c2 100644 --- a/src/everest/detached/client.py +++ b/src/everest/detached/client.py @@ -1,38 +1,27 @@ +from __future__ import annotations + import asyncio import logging -import re -import ssl -import time import traceback -from base64 import b64encode from collections.abc import Callable +from contextlib import closing from pathlib import Path -from typing import Any - -import requests -from pydantic import ValidationError -from websockets import ConnectionClosedError, ConnectionClosedOK -from websockets.sync.client import connect +from typing import TYPE_CHECKING, Any -from ert.run_models.event import EverestBatchResultEvent, status_event_from_json from ert.scheduler import create_driver from ert.scheduler.driver import Driver, FailedSubmit from ert.scheduler.event import StartedEvent -from ert.services.ert_client import ErtClient +from ert.services import ErtClient from ert.trace import get_traceparent -from everest.config import EverestConfig, ServerConfig +from everest.config import EverestConfig from everest.strings import ( OPT_PROGRESS_ID, SIM_PROGRESS_ID, - EverEndpoints, ) -# Specifies how many times to try a http request within the specified timeout. -_HTTP_REQUEST_RETRY = 10 +if TYPE_CHECKING: + from ert.run_models.event import EverestBatchResultEvent -# Proxy configuration for outgoing requests. -# For internal LAN HTTP requests not using a proxy is recommended. -PROXY = {"http": None, "https": None} # The methods in this file are typically called for the client side. # Information from the client side is relatively uninteresting, so we show it in @@ -70,144 +59,6 @@ async def start_server(config: EverestConfig, logging_level: int) -> Driver: return driver -def stop_server( - server_context: tuple[str, str, tuple[str, str]], retries: int = 5 -) -> bool: - """Stop server if found and it is running.""" - url, cert, auth = server_context - for retry in range(retries): - try: - stop_endpoint = f"{url}/{EverEndpoints.STOP}" - response = requests.post( - stop_endpoint, - verify=cert, - auth=auth, - proxies=PROXY, # type: ignore - ) - response.raise_for_status() - except Exception: - logger.debug(traceback.format_exc()) - time.sleep(retry) - else: - return True - return False - - -def get_experiments( - server_context: tuple[str, str, tuple[str, str]], - retries: int = 5, -) -> list[str]: - url, cert, auth = server_context - for retry in range(retries): - try: - response = requests.get( - f"{url}/{EverEndpoints.EXPERIMENTS}", - verify=cert, - auth=auth, - proxies=PROXY, - ) - response.raise_for_status() - return response.json()["experiment_ids"] - except Exception: - logger.debug(traceback.format_exc()) - time.sleep(retry) - raise RuntimeError("Failed to get experiment_ids") - - -def start_experiment( - server_context: tuple[str, str, tuple[str, str]], - config: EverestConfig, - retries: int = 5, -) -> str: - url, cert, auth = server_context - for retry in range(retries): - try: - start_endpoint = f"{url}/{EverEndpoints.START_EXPERIMENT}" - response = requests.post( - start_endpoint, - verify=cert, - auth=auth, - proxies=PROXY, # type: ignore - json=config.to_dict(), - ) - response.raise_for_status() - return response.json()["experiment_id"] - except Exception: - logger.debug(traceback.format_exc()) - time.sleep(retry) - raise RuntimeError("Failed to start experiment") - - -def extract_errors_from_file(path: str) -> list[str]: - return re.findall(r"(Error \w+.*)", Path(path).read_text(encoding="utf-8")) - - -def wait_for_server(api: ErtClient, timeout: float) -> None: - """ - Waits until the everest server has started. Polls - for server availability until timeout (measured in seconds). - - Timeout is not strict as the server status is polled periodically and - each underlying HTTP request has its own timeout, the wall-clock - duration may exceed the requested timeout slightly. - - Raises an exception if no response within the timeout. - """ - wait_start_time: float = time.monotonic() - while time.monotonic() - wait_start_time <= timeout: - if server_is_running( - *ServerConfig.get_server_context_from_conn_info(api.conn_info) - ): - return - until_timeout = max(0, timeout - (time.monotonic() - wait_start_time)) - time.sleep(min(1, until_timeout)) - raise RuntimeError( - "Failed to get reply from server " - f"within {time.monotonic() - wait_start_time:g} seconds" - ) - - -def wait_for_server_to_stop( - server_context: tuple[str, str, tuple[str, str]], timeout: int -) -> None: - """ - Checks everest server has stopped _HTTP_REQUEST_RETRY times. Waits - progressively longer between each check. - - Raise an exception when the timeout is reached. - """ - if server_is_running(*server_context): - sleep_time_increment = float(timeout) / (2**_HTTP_REQUEST_RETRY - 1) - for retry_count in range(_HTTP_REQUEST_RETRY): - sleep_time = sleep_time_increment * (2**retry_count) - time.sleep(sleep_time) - if not server_is_running(*server_context): - return - - # If number of retries reached and server still running - throw exception - if server_is_running(*server_context): - raise Exception("Failed to stop server within configured timeout.") - - -def server_is_running(url: str, cert: str, auth: tuple[str, str]) -> bool: - try: - logger.debug(f"Checking server status at {url} ") - if "None:None" in url: - return False - response = requests.get( - url, - verify=cert, - auth=auth, - timeout=1, - proxies=PROXY, # type: ignore - ) - response.raise_for_status() - except Exception: - logger.debug(traceback.format_exc()) - return False - return True - - def get_opt_status_from_batch_result_event( event: EverestBatchResultEvent, ) -> dict[str, Any]: @@ -228,8 +79,8 @@ def get_opt_status_from_batch_result_event( def start_monitor( - server_context: tuple[str, str, tuple[str, str]], - callback: Callable[..., None], + client: ErtClient, + callback: Callable[[dict[str, Any]], None], experiment_id: str, polling_interval: float = 0.1, ) -> None: @@ -238,46 +89,20 @@ def start_monitor( Monitoring stops when the server stops answering. """ - url, cert, auth = server_context - ssl_context = ssl.create_default_context() - ssl_context.load_verify_locations(cafile=cert) - username, password = auth - credentials = b64encode(f"{username}:{password}".encode()).decode() - - try: # ruff: ignore[too-many-statements-in-try-clause] - with connect( - url.replace("https://", "wss://") - + f"/{EverEndpoints.EVENTS}/{experiment_id}", - ssl=ssl_context, - open_timeout=30, - additional_headers={"Authorization": f"Basic {credentials}"}, - ) as websocket: - while True: - try: # ruff: ignore[too-many-statements-in-try-clause] - message = websocket.recv(timeout=1.0) - event = status_event_from_json(message) - if isinstance(event, EverestBatchResultEvent): - callback( - { - OPT_PROGRESS_ID: get_opt_status_from_batch_result_event( - event - ) - } - ) - else: - callback({SIM_PROGRESS_ID: event}) - except TimeoutError: - pass - except ConnectionClosedOK: - logger.debug("Connection closed") - break - except ConnectionClosedError: - logger.debug("Connection closed") - break - except ValidationError as e: - logger.error("Error when processing event %s", exc_info=e) - - time.sleep(polling_interval) + from ert.run_models.event import ( # ruff: ignore[import-outside-top-level] + EverestBatchResultEvent, + ) + try: + with closing( + client.iter_events(experiment_id, refresh_interval=polling_interval) + ) as events: + for event in events: + if isinstance(event, EverestBatchResultEvent): + callback( + {OPT_PROGRESS_ID: get_opt_status_from_batch_result_event(event)} + ) + else: + callback({SIM_PROGRESS_ID: event}) except Exception: logger.exception(traceback.format_exc()) diff --git a/src/everest/detached/everserver.py b/src/everest/detached/everserver.py index 0e33f79eb50..3a1db582294 100644 --- a/src/everest/detached/everserver.py +++ b/src/everest/detached/everserver.py @@ -19,7 +19,6 @@ from ert.trace import tracer from ert.utils import makedirs_if_needed from everest.config import ServerConfig -from everest.detached import get_experiments from everest.strings import ( DEFAULT_LOGGING_FORMAT, OPTIMIZATION_LOG_DIR, @@ -151,9 +150,7 @@ def main() -> None: client = ErtClient.get_client(Path(server_path)) done = False while not done: - experiment_ids = get_experiments( - ServerConfig.get_server_context_from_conn_info(client.conn_info) - ) + experiment_ids = client.experiment_ids() active = [ ExperimentStatus( **client.experiment_status(experiment_id) diff --git a/src/everest/gui/main_window.py b/src/everest/gui/main_window.py index 880d67422fa..dabbf35d3d1 100644 --- a/src/everest/gui/main_window.py +++ b/src/everest/gui/main_window.py @@ -1,6 +1,5 @@ from __future__ import annotations -import ssl from pathlib import Path from PyQt6.QtCore import pyqtSignal as Signal @@ -10,13 +9,13 @@ QMainWindow, ) +from ert.ensemble_evaluator.config import EvaluatorServerConfig from ert.gui.ertnotifier import ErtNotifier from ert.gui.experiments import RunDialog -from ert.gui.experiments.experiment_client import ExperimentClient from ert.plugins import ErtPluginManager +from ert.run_models.run_model import RunModelAPI from ert.services.ert_client import ErtClient from everest.config import ServerConfig -from everest.detached import get_experiments, wait_for_server class EverestMainWindow(QMainWindow): @@ -46,39 +45,37 @@ def run(self) -> None: client = ErtClient.get_client( Path(ServerConfig.get_session_dir(self.output_dir)) ) - wait_for_server(client, 60) - - server_context = ServerConfig.get_server_context_from_conn_info( - client.conn_info - ) - url, cert, auth = server_context - - ssl_context = ssl.create_default_context() - ssl_context.load_verify_locations(cafile=cert) - username, password = auth - - exp_client = ExperimentClient( - experiment_id=get_experiments(server_context)[-1], - url=url, - cert_file=cert, - username=username, - password=password, - ssl_context=ssl_context, + client.wait_for_server(60) + + experiment_id = client.experiment_ids()[-1] + config = client.experiment_config(experiment_id) + + config_filename = Path(config["config_path"]).name + self.setWindowTitle(f"EVEREST - {config_filename}") + + def start_fn( + evaluator_server_config: EvaluatorServerConfig, + *, + rerun_failed_realizations: bool = False, + ) -> None: + pass + + run_model_api = RunModelAPI( + experiment_name=config_filename, + supports_rerunning_failed_realizations=False, + start_simulations_thread=start_fn, + cancel=client.stop_experiment_server, # type: ignore + has_failed_realizations=lambda: False, ) - - config = exp_client.config - title = Path(config["config_path"]).name - self.setWindowTitle(f"EVEREST - {title}") - - run_model_api = exp_client.create_run_model_api() - event_queue, event_monitor_thread = ( - exp_client.setup_event_queue_from_ws_endpoint( - refresh_interval=0.02, open_timeout=40, websocket_recv_timeout=1.0 - ) + event_queue, event_monitor_thread = client.setup_event_queue_from_ws_endpoint( + experiment_id=experiment_id, + refresh_interval=0.02, + open_timeout=40, + websocket_recv_timeout=1.0, ) run_dialog = RunDialog( - title=title, + title=config_filename, run_model_api=run_model_api, event_queue=event_queue, notifier=ErtNotifier(), diff --git a/src/everest/strings.py b/src/everest/strings.py index 587f6a7b311..a5f7e0de5ba 100644 --- a/src/everest/strings.py +++ b/src/everest/strings.py @@ -1,5 +1,3 @@ -from enum import StrEnum, auto - DEFAULT_OUTPUT_DIR = "everest_output" DEFAULT_LOGGING_FORMAT = "%(asctime)s %(name)s %(levelname)s: %(message)s" @@ -17,14 +15,3 @@ SESSION_DIR = ".session" SIM_PROGRESS_ID = "simulation_progress" STORAGE_DIR = "simulation_results" - - -class EverEndpoints(StrEnum): - STOP = auto() - START_EXPERIMENT = auto() - CONFIG_PATH = auto() - START_TIME = auto() - EXPERIMENTS = auto() - STATUS = auto() - EVENTS = auto() - RUNPATH = auto() diff --git a/tests/ert/unit_tests/services/test_ert_client.py b/tests/ert/unit_tests/services/test_ert_client.py index 3f7c392f324..5ce268eddd1 100644 --- a/tests/ert/unit_tests/services/test_ert_client.py +++ b/tests/ert/unit_tests/services/test_ert_client.py @@ -1,10 +1,14 @@ import io from typing import Any +from unittest.mock import MagicMock +import httpx import pandas as pd import pytest +from ert.ensemble_evaluator import EndEvent from ert.services.ert_client import ErtClient +from ert.services.shared_client import SharedClient class RecordingResponse: @@ -126,3 +130,54 @@ def test_that_mutating_a_returned_parameter_frame_leaves_the_cache_intact(api, c # Assert that the cached value is not mutated. assert api.parameter("ens_1", "gen_kw")["0"].to_list() == [1.0, 2.0, 3.0] + + +@pytest.fixture +def event_client(monkeypatch): + transport = MagicMock(spec=SharedClient) + transport.conn_info.base_url = "https://localhost:1234" + transport.conn_info.auth_token = "token" + transport.conn_info.cert = False + connection = MagicMock() + connect = MagicMock() + connect.return_value.__enter__.return_value = connection + monkeypatch.setattr("ert.services.ert_client.connect", connect) + return ErtClient(transport), connection, connect + + +def test_that_closing_event_iterator_releases_connection(event_client): + api, connection, connect = event_client + event = EndEvent(failed=False, msg="completed") + connection.recv.return_value = event.model_dump_json() + events = api.iter_events("experiment") + + assert next(events) == event + connect.return_value.__exit__.assert_not_called() + events.close() + connect.return_value.__exit__.assert_called_once() + + +@pytest.mark.parametrize("status_code", [200, 401, 503]) +@pytest.mark.parametrize("timeout", [None, 1.0]) +def test_that_server_probe_checks_authenticated_endpoint_with_requested_timeout( + event_client, status_code, timeout +): + api, _, _ = event_client + api.client.request.return_value.status_code = status_code + + assert api.server_is_running(timeout=timeout) is (status_code == 200) + + api.client.request.assert_called_once_with( + "GET", + "/experiment_server/", + auth=("username", "token"), + timeout=120 if timeout is None else timeout, + ) + + +@pytest.mark.parametrize("error_type", [httpx.ConnectError, httpx.ReadTimeout]) +def test_that_server_probe_returns_false_on_transport_error(event_client, error_type): + api, _, _ = event_client + api.client.request.side_effect = error_type("unavailable") + + assert not api.server_is_running(timeout=1) diff --git a/tests/everest/entry_points/test_everest_entry.py b/tests/everest/entry_points/test_everest_entry.py index cd045b6e3ec..c06e4052071 100644 --- a/tests/everest/entry_points/test_everest_entry.py +++ b/tests/everest/entry_points/test_everest_entry.py @@ -1,11 +1,10 @@ import logging import tempfile from pathlib import Path -from unittest.mock import MagicMock, patch +from unittest.mock import DEFAULT, patch import pytest -import everest from ert.config import QueueSystem from ert.run_models.everest_run_model import ExperimentStatus from ert.storage import ExperimentState @@ -22,20 +21,16 @@ def raise_system_error(*args, **kwargs): @patch("everest.bin.everest_script.run_detached_monitor") -@patch("everest.bin.everest_script.wait_for_server") @patch("everest.bin.everest_script.start_server") @patch("everest.config.ServerConfig.get_server_context_from_conn_info") @patch( "everest.bin.everest_script.ErtClient", - **{"get_client.side_effect": [TimeoutError(), MagicMock()]}, + **{"get_client.side_effect": [TimeoutError(), DEFAULT]}, ) -@patch("everest.bin.everest_script.start_experiment") def test_everest_entry_debug( - start_experiment_mock, everest_script_api_mock, get_server_context_from_conn_info_mock, start_server_mock, - wait_for_server_mock, start_monitor_mock, caplog, change_to_tmpdir, @@ -57,11 +52,15 @@ def test_everest_entry_debug( everest_entry(["config.yml", "--debug"]) logstream = "\n".join(caplog.messages) start_server_mock.assert_called_once() - wait_for_server_mock.assert_called_once() - start_monitor_mock.assert_called_once() - start_experiment_mock.assert_called_once() + everest_script_api_mock.get_client.return_value.wait_for_server.assert_called_once_with( + timeout=600 + ) + start_monitor_mock.assert_called_once_with( + client=everest_script_api_mock.get_client.return_value, + experiment_id=everest_script_api_mock.get_client.return_value.start_experiment.return_value, + ) assert everest_script_api_mock.get_client.call_count == 2 - assert get_server_context_from_conn_info_mock.call_count == 2 + get_server_context_from_conn_info_mock.assert_not_called() # the config file itself is dumped at DEBUG level assert '"controls"' in logstream @@ -71,20 +70,16 @@ def test_everest_entry_debug( @patch("everest.bin.everest_script.run_detached_monitor") -@patch("everest.bin.everest_script.wait_for_server") @patch("everest.bin.everest_script.start_server") @patch("everest.config.ServerConfig.get_server_context_from_conn_info") @patch( "everest.bin.everest_script.ErtClient", - **{"get_client.side_effect": [TimeoutError(), MagicMock()]}, + **{"get_client.side_effect": [TimeoutError(), DEFAULT]}, ) -@patch("everest.bin.everest_script.start_experiment") def test_everest_entry( - start_experiment_mock, everest_script_api_mock, get_server_context_from_conn_info_mock, start_server_mock, - wait_for_server_mock, start_monitor_mock, change_to_tmpdir, ): @@ -95,26 +90,31 @@ def test_everest_entry( config.write_to_file("config.yml") everest_entry(["config.yml"]) start_server_mock.assert_called_once() - wait_for_server_mock.assert_called_once() - start_monitor_mock.assert_called_once() - start_experiment_mock.assert_called_once() + everest_script_api_mock.get_client.return_value.start_experiment.assert_called_once_with( + start_server_mock.call_args.args[0].to_dict() + ) + everest_script_api_mock.get_client.return_value.wait_for_server.assert_called_once_with( + timeout=600 + ) + start_monitor_mock.assert_called_once_with( + client=everest_script_api_mock.get_client.return_value, + experiment_id=everest_script_api_mock.get_client.return_value.start_experiment.return_value, + ) assert everest_script_api_mock.get_client.call_count == 2 - assert get_server_context_from_conn_info_mock.call_count == 2 + get_server_context_from_conn_info_mock.assert_not_called() @patch("everest.bin.everest_script.run_detached_monitor") -@patch("everest.bin.everest_script.wait_for_server") @patch("everest.bin.everest_script.start_server") -@patch("everest.bin.everest_script.start_experiment") @patch("everest.config.ServerConfig.get_server_context_from_conn_info") @patch( "everest.bin.everest_script.ErtClient", **{ "get_client.side_effect": [ TimeoutError(), - MagicMock(), + DEFAULT, TimeoutError(), - MagicMock(), + DEFAULT, ] }, ) @@ -126,9 +126,7 @@ def test_everest_entry_detached_already_run( kill_script_api_mock, everest_script_api_mock, get_server_context_from_conn_info_mock, - start_experiment_mock, start_server_mock, - wait_for_server_mock, start_monitor_mock, change_to_tmpdir, ): @@ -139,6 +137,9 @@ def test_everest_entry_detached_already_run( Path("config.yml").touch() config = everest_config_with_defaults(config_path="./config.yml") config.write_to_file("config.yml") + start_experiment_mock = ( + everest_script_api_mock.get_client.return_value.start_experiment + ) # start a new run everest_entry(["config.yml"]) @@ -198,17 +199,14 @@ def test_everest_entry_detached_already_run_monitor( @patch("everest.bin.everest_script.ErtClient") @patch("everest.config.ServerConfig.get_server_context_from_conn_info") @patch("everest.bin.everest_script.run_detached_monitor") -@patch("everest.bin.everest_script.wait_for_server") @patch("everest.bin.everest_script.start_server") -@patch("everest.bin.kill_script.stop_server", return_value=True) -@patch("everest.bin.kill_script.wait_for_server_to_stop") -@patch("everest.bin.kill_script.ErtClient") +@patch( + "everest.bin.kill_script.ErtClient", + **{"get_client.return_value.stop_experiment_server.return_value": True}, +) def test_everest_entry_detached_running( kill_api_mock, - wait_for_server_to_stop_mock, - stop_server_mock, start_server_mock, - wait_for_server_mock, start_monitor_mock, get_server_context_from_conn_info_mock, everest_script_api_mock, @@ -219,6 +217,7 @@ def test_everest_entry_detached_running( Path("config.yml").touch() config = everest_config_with_defaults(config_path="./config.yml") config.write_to_file("config.yml") + stop_server_mock = kill_api_mock.get_client.return_value.stop_experiment_server # can't start a new run if one is already running with capture_streams() as (out, _): @@ -227,7 +226,7 @@ def test_everest_entry_detached_running( assert "everest monitor" in out.getvalue() start_server_mock.assert_not_called() start_monitor_mock.assert_not_called() - wait_for_server_mock.assert_not_called() + everest_script_api_mock.get_client.return_value.wait_for_server.assert_not_called() everest_script_api_mock.get_client.assert_called_once() everest_script_api_mock.reset_mock() get_server_context_from_conn_info_mock.assert_not_called() @@ -235,11 +234,13 @@ def test_everest_entry_detached_running( # stop the server kill_entry(["config.yml"]) stop_server_mock.assert_called_once() - wait_for_server_to_stop_mock.assert_called_once() + kill_api_mock.get_client.return_value.wait_for_server_to_stop.assert_called_once_with( + timeout=60 + ) kill_api_mock.get_client.assert_called_once() kill_api_mock.reset_mock() - get_server_context_from_conn_info_mock.assert_called_once() - wait_for_server_mock.assert_not_called() + get_server_context_from_conn_info_mock.assert_not_called() + everest_script_api_mock.get_client.return_value.wait_for_server.assert_not_called() # if already running, nothing happens assert "everest kill" in out.getvalue() @@ -252,12 +253,11 @@ def test_everest_entry_detached_running( @patch("everest.bin.monitor_script.run_detached_monitor") @patch("everest.config.ServerConfig.get_server_context_from_conn_info") @patch( - "everest.bin.monitor_script.get_experiments", return_value=["test-experiment-id"] + "everest.bin.monitor_script.ErtClient", + **{"get_client.return_value.experiment_ids.return_value": ["test-experiment-id"]}, ) -@patch("everest.bin.monitor_script.ErtClient") def test_everest_entry_detached_running_monitor( monitor_script_api_mock, - get_experiments_mock, get_server_context_from_conn_info_mock, start_monitor_mock, change_to_tmpdir, @@ -271,10 +271,13 @@ def test_everest_entry_detached_running_monitor( # Attach to a running optimization. with capture_streams(): monitor_entry(["config.yml"]) - start_monitor_mock.assert_called_once() + start_monitor_mock.assert_called_once_with( + client=monitor_script_api_mock.get_client.return_value, + experiment_id="test-experiment-id", + ) monitor_script_api_mock.get_client.assert_called_once() - get_server_context_from_conn_info_mock.assert_called_once() - get_experiments_mock.assert_called_once() + get_server_context_from_conn_info_mock.assert_not_called() + monitor_script_api_mock.get_client.return_value.experiment_ids.assert_called_once() @patch("everest.bin.monitor_script.run_detached_monitor") @@ -307,29 +310,20 @@ def test_everest_entry_monitor_already_run( get_server_context_from_conn_info_mock.assert_not_called() -@pytest.fixture(autouse=True) -def mock_ssl(monkeypatch): - monkeypatch.setattr(everest.detached.client, "ssl", MagicMock()) - - @patch( "everest.bin.everest_script.run_detached_monitor", side_effect=raise_system_error, ) -@patch("everest.bin.everest_script.wait_for_server") @patch("everest.bin.everest_script.start_server") -@patch("everest.bin.everest_script.start_experiment") @patch("everest.config.ServerConfig.get_server_context_from_conn_info") @patch( "everest.bin.everest_script.ErtClient", - **{"get_client.side_effect": [TimeoutError(), MagicMock()]}, + **{"get_client.side_effect": [TimeoutError(), DEFAULT]}, ) def test_exception_raised_when_server_run_fails( everest_script_api_mock, get_server_context_from_conn_info_mock, - start_experiment_mock, start_server_mock, - wait_for_server_mock, start_monitor_mock, change_to_tmpdir, ): @@ -347,12 +341,11 @@ def test_exception_raised_when_server_run_fails( ) @patch("everest.config.ServerConfig.get_server_context_from_conn_info") @patch( - "everest.bin.monitor_script.get_experiments", return_value=["test-experiment-id"] + "everest.bin.monitor_script.ErtClient", + **{"get_client.return_value.experiment_ids.return_value": ["test-experiment-id"]}, ) -@patch("everest.bin.monitor_script.ErtClient") def test_exception_raised_when_server_run_fails_monitor( monitor_script_api_mock, - get_experiments_mock, get_server_context_from_conn_info_mock, start_monitor_mock, change_to_tmpdir, @@ -418,15 +411,13 @@ def test_that_run_everest_prints_where_it_runs( with ( patch( "everest.bin.everest_script.ErtClient", - **{"get_client.side_effect": [TimeoutError(), MagicMock()]}, + **{"get_client.side_effect": [TimeoutError(), DEFAULT]}, ), patch( "everest.config.ServerConfig.get_server_context_from_conn_info", return_value=("a", "b", ("c", "d")), ), patch("everest.bin.everest_script.start_server"), - patch("everest.bin.everest_script.wait_for_server"), - patch("everest.bin.everest_script.start_experiment"), ): everest_entry(["config.yml"]) diff --git a/tests/everest/entry_points/test_everexport.py b/tests/everest/entry_points/test_everexport.py deleted file mode 100644 index e69de29bb2d..00000000000 diff --git a/tests/everest/test_detached.py b/tests/everest/test_detached.py index d9eb010e379..14df3379b82 100644 --- a/tests/everest/test_detached.py +++ b/tests/everest/test_detached.py @@ -32,12 +32,7 @@ from everest.config.server_config import ServerConfig from everest.config.simulator_config import SimulatorConfig from everest.detached import ( - PROXY, - server_is_running, start_server, - stop_server, - wait_for_server, - wait_for_server_to_stop, ) from tests.everest.utils import everest_config_with_defaults @@ -47,6 +42,7 @@ @pytest.mark.xdist_group(name="starts_everest") @pytest.mark.usefixtures("use_site_configurations_with_no_queue_options") async def test_https_requests(change_to_tmpdir): + proxies = {"http": None, "https": None} Path("./config.yml").touch() everest_config = everest_config_with_defaults(config_path="./config.yml") everest_config.forward_model.append(ForwardModelStepConfig(job="sleep 5")) @@ -63,44 +59,99 @@ async def test_https_requests(change_to_tmpdir): client = ErtClient.get_client( Path(ServerConfig.get_session_dir(everest_config.output_dir)), 240 ) - wait_for_server(client, 240) + client.wait_for_server(240) url, cert, auth = ServerConfig.get_server_context_from_conn_info(client.conn_info) - result = requests.get(url, verify=cert, auth=auth, proxies=PROXY) # ruff: ignore[blocking-http-call-in-async-function] + result = requests.get(url, verify=cert, auth=auth, proxies=proxies) # ruff: ignore[blocking-http-call-in-async-function] assert result.status_code == 200 # Request has succeeded # Test http request fail http_url = url.replace("https", "http") with pytest.raises(Exception): # ruff: ignore[assert-raises-exception, pytest-raises-too-broad, pytest-raises-with-multiple-statements] B017 - response = requests.get(http_url, verify=cert, auth=auth, proxies=PROXY) # ruff: ignore[blocking-http-call-in-async-function] + response = requests.get(http_url, verify=cert, auth=auth, proxies=proxies) # ruff: ignore[blocking-http-call-in-async-function] response.raise_for_status() # Test request with wrong password fails auth = ("admin", "wrong_password") - result = requests.get(url, verify=cert, auth=auth, proxies=PROXY) # ruff: ignore[blocking-http-call-in-async-function] + result = requests.get(url, verify=cert, auth=auth, proxies=proxies) # ruff: ignore[blocking-http-call-in-async-function] assert result.status_code == 401 # Unauthorized # Test stopping server - assert server_is_running( - *ServerConfig.get_server_context_from_conn_info(client.conn_info) - ) - server_context = ServerConfig.get_server_context_from_conn_info(client.conn_info) - if stop_server(server_context): - wait_for_server_to_stop(server_context, 240) - assert not server_is_running(*server_context) + assert client.server_is_running(timeout=1) + if client.stop_experiment_server(): + client.wait_for_server_to_stop(240) + assert not client.server_is_running(timeout=1) -@patch("everest.detached.server_is_running", return_value=False) -@patch( - "everest.config.ServerConfig.get_server_context_from_conn_info", - return_value=("url", "cert", ("user", "token")), -) -def test_wait_for_server(mock_get_context, mock_is_running): - client = MagicMock() +@pytest.fixture +def polling_client(monkeypatch): + client = ErtClient(MagicMock()) + monkeypatch.setattr(client, "server_is_running", MagicMock()) + return client + + +@patch("ert.services.ert_client.time") +def test_that_wait_for_server_raises_when_server_remains_unavailable( + clock, polling_client +): + client = polling_client + client.server_is_running.return_value = False + clock.monotonic.side_effect = [0, 0, 0, 2, 2] with pytest.raises( RuntimeError, match=r"Failed to get reply from server within .* seconds" ): - wait_for_server(client, timeout=0.01) + client.wait_for_server(timeout=1) + + client.server_is_running.assert_called_once_with(timeout=1) + clock.sleep.assert_called_once_with(1) + + +@pytest.mark.parametrize("states", [[True], [False, True]]) +@patch("ert.services.ert_client.time") +def test_that_wait_for_server_returns_when_server_becomes_available( + clock, states, polling_client +): + client = polling_client + client.server_is_running.side_effect = states + clock.monotonic.return_value = 0 + + client.wait_for_server(timeout=10) + + assert client.server_is_running.call_args_list == [mock.call(timeout=1)] * len( + states + ) + assert clock.sleep.call_count == len(states) - 1 + + +@pytest.mark.parametrize("states", [[False, False], [True, False], [True, True, False]]) +@patch("ert.services.ert_client.time") +def test_that_wait_for_server_to_stop_returns_when_server_is_unavailable( + clock, states, polling_client +): + client = polling_client + client.server_is_running.side_effect = states + + client.wait_for_server_to_stop(timeout=10) + + assert client.server_is_running.call_args_list == [mock.call(timeout=1)] * len( + states + ) + assert clock.sleep.call_count == (len(states) - 1 if states[0] else 0) + + +@patch("ert.services.ert_client.time") +def test_that_wait_for_server_to_stop_raises_after_all_retries(clock, polling_client): + client = polling_client + client.server_is_running.return_value = True + + with pytest.raises( + Exception, match="Failed to stop server within configured timeout" + ): + client.wait_for_server_to_stop(timeout=10) + + assert client.server_is_running.call_args_list == [mock.call(timeout=1)] * 12 + assert clock.sleep.call_count == 10 + assert sum(call.args[0] for call in clock.sleep.call_args_list) == pytest.approx(10) @pytest.mark.usefixtures("use_site_configurations_with_no_queue_options") diff --git a/tests/everest/test_everest_client.py b/tests/everest/test_everest_client.py deleted file mode 100644 index dbcb9e69a08..00000000000 --- a/tests/everest/test_everest_client.py +++ /dev/null @@ -1,266 +0,0 @@ -import logging -import ssl -import sys -import threading -import warnings -from pathlib import Path - -import pytest -import requests -import uvicorn -import yaml -from fastapi import FastAPI -from starlette.responses import Response - -from ert.gui.experiments.experiment_client import ExperimentClient -from ert.run_models.event import EverestBatchResultEvent, EverestStatusEvent -from ert.services import ErtClient, SharedClient -from ert.shared import find_available_socket -from everest.bin.everest_script import everest_entry -from everest.config import EverestConfig, ServerConfig -from everest.detached import get_experiments, server_is_running -from everest.strings import EverEndpoints -from tests.ert.utils import wait_until - - -@pytest.fixture -def client_server_mock() -> tuple[FastAPI, threading.Thread, ExperimentClient]: - server_app = FastAPI() - host = "127.0.0.1" - port = find_available_socket(host, range(5000, 5800)).getsockname()[1] - server_url = f"http://{host}:{port}" - - @server_app.get("alive") - def alive(): - return Response("Hello", status_code=200) - - server = uvicorn.Server( - uvicorn.Config(server_app, host=host, port=port, log_level="info") - ) - - server_thread = threading.Thread( - target=server.run, - daemon=True, - ) - - everest_client = ExperimentClient( - experiment_id="test_experiment_id", - url=server_url, - cert_file="N/A", - username="", - password="", - ssl_context=ssl.create_default_context(), - ) - - def wait_until_alive(timeout=60, sleep_between_retries=1) -> None: - def ping_server() -> bool: - try: - requests.get( - f"{server_url}/alive", - verify="N/A", - auth=("", ""), - proxies={"http": None, "https": None}, # type: ignore - ) - except requests.exceptions.ConnectionError: - return False - else: - return True - - # These warnings emitted by uvicorn, which is still using legacy - # websockets. This is a known issue, and does not cause problems in the - # main code. (see: https://github.com/encode/uvicorn/discussions/2476) - # Hence we ignore them in the tests. Potentially, this may be - # removed when this is resolved within uvicorn. - with warnings.catch_warnings(): - warnings.filterwarnings("ignore", message="websockets.legacy is deprecated") - warnings.filterwarnings( - "ignore", - message="websockets.server.WebSocketServerProtocol is deprecated", - ) - wait_until(ping_server, timeout=timeout, interval=sleep_between_retries) - - yield server_app, server_thread, everest_client, wait_until_alive - - if server_thread.is_alive(): - server.should_exit = True - server_thread.join() - - -@pytest.mark.slow -@pytest.mark.flaky(rerun=2) -def test_that_stop_invokes_correct_endpoint( - caplog, client_server_mock: tuple[FastAPI, threading.Thread, ExperimentClient] -): - server_app, server_thread, client, wait_until_alive = client_server_mock - - @server_app.post(f"/{EverEndpoints.STOP}") - def stop(): - return Response("STOP..", 200) - - server_thread.start() - wait_until_alive() - - with caplog.at_level(logging.INFO): - client.stop() - - assert "Cancelled experiment from EVEREST" in caplog.messages - server_thread.should_exit = True - - -@pytest.mark.slow -def test_that_stop_errors_on_non_ok_httpcode( - caplog, client_server_mock: tuple[FastAPI, threading.Thread, ExperimentClient] -): - server_app, server_thread, client, wait_until_alive = client_server_mock - - @server_app.post(f"/{EverEndpoints.STOP}") - def stop(): - return Response("STOP..", 505) - - server_thread.start() - wait_until_alive() - - with caplog.at_level(logging.ERROR): - client.stop() - - assert any( - "Failed to cancel EVEREST experiment" in m - and "server responded with status 505" in m - for m in caplog.messages - ) - - -def test_that_stop_errors_on_server_down( - caplog, client_server_mock: tuple[FastAPI, threading.Thread, ExperimentClient] -): - _, _, client, _ = client_server_mock - - with caplog.at_level(logging.ERROR): - client.stop() - - assert any( - "Connection error when cancelling EVEREST experiment" in m - for m in caplog.messages - ) - - -@pytest.mark.slow -def test_that_stop_errors_on_server_up_but_endpoint_down( - caplog, client_server_mock: tuple[FastAPI, threading.Thread, ExperimentClient] -): - _, server_thread, client, wait_until_alive = client_server_mock - - server_thread.start() - wait_until_alive() - - with caplog.at_level(logging.ERROR): - client.stop() - - assert any( - "server responded with status 404: Not Found" in m for m in caplog.messages - ) - - -@pytest.mark.skip_mac_ci -@pytest.mark.slow -@pytest.mark.xdist_group("math_func/config_minimal.yml") -@pytest.mark.flaky(rerun=3) -@pytest.mark.skipif( - sys.version_info[0:3] == (3, 13, 6), reason="Fails on Python 3.13.6" -) -def test_that_multiple_everest_clients_can_connect_to_server( - cached_example, change_to_tmpdir -): - # We use a cached run for the reference list of received events - path, config_file, _, server_events_list = cached_example( - "math_func/config_minimal.yml" - ) - SharedClient.close_client() - - config_path = Path(path) / config_file - config_content = yaml.safe_load(config_path.read_text(encoding="utf-8")) - config_content["simulator"] = {"queue_system": {"name": "local", "max_running": 2}} - config_path.write_text( - yaml.dump(config_content, default_flow_style=False), encoding="utf-8" - ) - - ever_config = EverestConfig.load_file(config_path) - - # Run the case through everserver - everest_main_thread = threading.Thread( - target=everest_entry, args=[[str(config_path)]] - ) - - everest_main_thread.start() - session_dir = Path(ServerConfig.get_session_dir(ever_config.output_dir)) - - def everserver_is_running() -> bool: - try: - api = ErtClient.get_client(session_dir, connect_timeout=1) - except TimeoutError: - return False - return server_is_running( - *ServerConfig.get_server_context_from_conn_info(api.conn_info) - ) - - wait_until(everserver_is_running, interval=1, timeout=300) - - api = ErtClient.get_client(session_dir) - server_context = ServerConfig.get_server_context_from_conn_info(api.conn_info) - url, cert, auth = server_context - - ssl_context = ssl.create_default_context() - ssl_context.load_verify_locations(cafile=cert) - username, password = auth - - client_event_queues = [] - monitor_threads = [] - for _ in range(5): - client = ExperimentClient( - experiment_id=get_experiments(server_context)[-1], - url=url, - cert_file=cert, - username=username, - password=password, - ssl_context=ssl_context, - ) - - # Connect to the websockets endpoint - client_event_queue, monitor_thread = client.setup_event_queue_from_ws_endpoint() - client_event_queues.append(client_event_queue) - monitor_threads.append(monitor_thread) - monitor_thread.start() - - # Wait until the server has finished running the simulation - everest_main_thread.join() - for _thread in monitor_threads: - if _thread.is_alive(): - _thread.join(timeout=5) - - # Expect all the clients to hold the same events - client_event_lists = [] - for event_queue in client_event_queues: - event_list = [] - while not event_queue.empty(): - event_list.append(event_queue.get()) - - client_event_lists.append(event_list) - - first = client_event_lists[0] - assert all(first == other for other in client_event_lists[1:]) - - everest_event_types = (EverestStatusEvent, EverestBatchResultEvent) - - first_everevents = [ - e.event_type for e in first if isinstance(e, everest_event_types) - ] - assert len(first_everevents) > 0 - - server_everevents = [ - e.event_type for e in server_events_list if isinstance(e, everest_event_types) - ] - assert len(server_everevents) > 0 - - # Compare only everest events, as the events from the forward model - # are (at time of writing) not deterministic enough to expect equality - assert first_everevents == server_everevents diff --git a/tests/everest/test_everest_output.py b/tests/everest/test_everest_output.py index 133eefc1545..3dbe7cd570b 100644 --- a/tests/everest/test_everest_output.py +++ b/tests/everest/test_everest_output.py @@ -32,13 +32,9 @@ def test_that_one_experiment_creates_one_ensemble_per_batch(cached_example): ) @patch("everest.config.ServerConfig.get_server_context_from_conn_info") @patch("everest.bin.everest_script.run_detached_monitor") -@patch("everest.bin.everest_script.wait_for_server") @patch("everest.bin.everest_script.start_server") -@patch("everest.bin.everest_script.start_experiment") def test_save_running_config( - mock_start_experiment, mock_start_server, - mock_wait_for_server, mock_run_detached_monitor, mock_get_server_context, mock_start_session, diff --git a/tests/everest/test_everserver.py b/tests/everest/test_everserver.py index 9fb8c4be8de..6ada9aa65fa 100644 --- a/tests/everest/test_everserver.py +++ b/tests/everest/test_everserver.py @@ -28,9 +28,7 @@ from everest.config import EverestConfig, ServerConfig from everest.detached import ( everserver, - start_experiment, start_server, - wait_for_server, ) from everest.strings import ( OPT_FAILURE_ALL_REALIZATIONS, @@ -71,11 +69,8 @@ async def server_running(): driver = await start_server(config, logging.DEBUG) api = ErtClient.get_client(Path(ServerConfig.get_session_dir(config.output_dir))) - wait_for_server(api, 120) - start_experiment( - server_context=ServerConfig.get_server_context_from_conn_info(api.conn_info), - config=config, - ) + api.wait_for_server(timeout=120) + api.start_experiment(config.to_dict()) await server_running() diff --git a/tests/everest/test_monitor.py b/tests/everest/test_monitor.py index f907e46d385..f5a74fca40d 100644 --- a/tests/everest/test_monitor.py +++ b/tests/everest/test_monitor.py @@ -4,14 +4,11 @@ import string from collections import defaultdict from datetime import UTC, datetime -from functools import partial -from unittest.mock import MagicMock, patch +from unittest.mock import MagicMock import pytest from fastapi.encoders import jsonable_encoder -from websockets.sync.client import ClientConnection -import everest from ert.ensemble_evaluator import ( EndEvent, FullSnapshotEvent, @@ -20,8 +17,11 @@ ) from ert.ensemble_evaluator.snapshot import EnsembleSnapshotMetadata from ert.resources import all_shell_script_fm_steps -from ert.run_models.event import EverestBatchResultEvent -from everest.bin.utils import run_detached_monitor +from ert.run_models.event import EverestBatchResultEvent, status_event_from_json +from ert.services import ErtClient +from everest.bin.utils import run_detached_monitor, run_empty_detached_monitor +from everest.detached.client import start_monitor +from everest.strings import SIM_PROGRESS_ID from tests.ert.utils import SnapshotBuilder METADATA = EnsembleSnapshotMetadata( @@ -39,6 +39,62 @@ def fixed_terminal_width(monkeypatch): ) +@pytest.fixture +def monitor_client(): + return MagicMock(spec=ErtClient) + + +def test_that_monitor_delivers_events_after_end_event(monitor_client): + events = [EndEvent(failed=False, msg="first"), EndEvent(failed=True, msg="last")] + monitor_client.iter_events.return_value = (event for event in events) + callback = MagicMock() + + start_monitor(monitor_client, callback, "experiment", polling_interval=0.2) + + monitor_client.iter_events.assert_called_once_with( + "experiment", refresh_interval=0.2 + ) + assert [call.args[0][SIM_PROGRESS_ID] for call in callback.call_args_list] == events + + +def test_that_empty_monitor_consumes_all_events_without_output(monitor_client, capsys): + consumed = [] + + def iter_events(): + for message in ["first", "last"]: + yield EndEvent(failed=False, msg=message) + consumed.append(message) + + monitor_client.iter_events.return_value = iter_events() + + run_empty_detached_monitor(monitor_client, "experiment") + + assert consumed == ["first", "last"] + assert not capsys.readouterr().out + + +@pytest.mark.parametrize("exception_type", [RuntimeError, KeyboardInterrupt]) +def test_that_callback_exception_closes_event_iterator(monitor_client, exception_type): + closed = MagicMock() + + def iter_events(): + try: + yield EndEvent(failed=False, msg="completed") + finally: + closed() + + monitor_client.iter_events.return_value = iter_events() + callback = MagicMock(side_effect=exception_type) + + if exception_type is KeyboardInterrupt: + with pytest.raises(KeyboardInterrupt): + start_monitor(monitor_client, callback, "experiment") + else: + start_monitor(monitor_client, callback, "experiment") + + closed.assert_called_once() + + @pytest.fixture def full_snapshot_event(): snapshot = SnapshotBuilder(metadata=METADATA) @@ -178,21 +234,18 @@ def snapshot_update_event_with_fm_message(): @pytest.mark.slow def test_that_the_monitor_shows_failed_jobs( - monkeypatch, full_snapshot_event, snapshot_update_failure_event, capsys + monitor_client, full_snapshot_event, snapshot_update_failure_event, capsys ): - server_mock = MagicMock() - connection_mock = MagicMock(spec=ClientConnection) - connection_mock.recv.side_effect = [ - full_snapshot_event, - snapshot_update_failure_event, - json.dumps(jsonable_encoder(EndEvent(failed=True, msg="Failed"))), - ] - server_mock.return_value.__enter__.return_value = connection_mock - monkeypatch.setattr(everest.detached.client, "connect", server_mock) - monkeypatch.setattr(everest.detached.client, "ssl", MagicMock()) - partial(everest.detached.start_monitor, polling_interval=0.1) + monitor_client.iter_events.return_value = ( + status_event_from_json(message) + for message in [ + full_snapshot_event, + snapshot_update_failure_event, + json.dumps(jsonable_encoder(EndEvent(failed=True, msg="Failed"))), + ] + ) run_detached_monitor( - ("some/url", "cert", ("username", "password")), + monitor_client, experiment_id="test-experiment-id", ) captured = capsys.readouterr() @@ -213,26 +266,19 @@ def test_that_the_monitor_shows_failed_jobs( @pytest.mark.slow def test_that_the_monitor_shows_running_jobs( - monkeypatch, full_snapshot_event, snapshot_update_event, capsys + monitor_client, full_snapshot_event, snapshot_update_event, capsys ): - server_mock = MagicMock() - connection_mock = MagicMock(spec=ClientConnection) - connection_mock.recv.side_effect = [ - full_snapshot_event, - snapshot_update_event, - json.dumps( - jsonable_encoder(EndEvent(failed=False, msg="Experiment completed")) - ), - ] - server_mock.return_value.__enter__.return_value = connection_mock - monkeypatch.setattr(everest.detached.client, "connect", server_mock) - monkeypatch.setattr(everest.detached.client, "ssl", MagicMock()) - patched = partial(everest.detached.start_monitor, polling_interval=0.1) - with patch("everest.bin.utils.start_monitor", patched): - run_detached_monitor( - ("some/url", "cert", ("username", "password")), - experiment_id="test-experiment-id", - ) + monitor_client.iter_events.return_value = ( + status_event_from_json(message) + for message in [ + full_snapshot_event, + snapshot_update_event, + json.dumps( + jsonable_encoder(EndEvent(failed=False, msg="Experiment completed")) + ), + ] + ) + run_detached_monitor(monitor_client, experiment_id="test-experiment-id") captured = capsys.readouterr() expected = [ "============ Running forward models (Batch #0) =============\n", @@ -248,21 +294,18 @@ def test_that_the_monitor_shows_running_jobs( @pytest.mark.slow def test_that_a_forward_model_message_reaches_the_cli( - monkeypatch, full_snapshot_event, snapshot_update_event_with_fm_message, capsys + monitor_client, full_snapshot_event, snapshot_update_event_with_fm_message, capsys ): - server_mock = MagicMock() - connection_mock = MagicMock(spec=ClientConnection) - connection_mock.recv.side_effect = [ - full_snapshot_event, - snapshot_update_event_with_fm_message, - json.dumps(jsonable_encoder(EndEvent(failed=True, msg="Failed"))), - ] - server_mock.return_value.__enter__.return_value = connection_mock - monkeypatch.setattr(everest.detached.client, "connect", server_mock) - monkeypatch.setattr(everest.detached.client, "ssl", MagicMock()) - partial(everest.detached.start_monitor, polling_interval=0.1) + monitor_client.iter_events.return_value = ( + status_event_from_json(message) + for message in [ + full_snapshot_event, + snapshot_update_event_with_fm_message, + json.dumps(jsonable_encoder(EndEvent(failed=True, msg="Failed"))), + ] + ) run_detached_monitor( - ("some/url", "cert", ("username", "password")), + monitor_client, experiment_id="test-experiment-id", ) captured = capsys.readouterr() @@ -288,22 +331,16 @@ def test_that_a_forward_model_message_reaches_the_cli( @pytest.mark.slow def test_that_a_failed_everest_batch_result_event_is_shown( - monkeypatch, everest_batch_result_event, capsys + monitor_client, everest_batch_result_event, capsys ): - server_mock = MagicMock() - connection_mock = MagicMock(spec=ClientConnection) - connection_mock.recv.side_effect = [ - everest_batch_result_event, - json.dumps(jsonable_encoder(EndEvent(failed=True, msg="Failed"))), - ] - server_mock.return_value.__enter__.return_value = connection_mock - monkeypatch.setattr(everest.detached.client, "connect", server_mock) - monkeypatch.setattr(everest.detached.client, "ssl", MagicMock()) - patched = partial(everest.detached.start_monitor, polling_interval=0.1) - with patch("everest.bin.utils.start_monitor", patched): - run_detached_monitor( - ("some/url", "cert", ("username", "password")), experiment_id="test-run-id" - ) + monitor_client.iter_events.return_value = ( + status_event_from_json(message) + for message in [ + everest_batch_result_event, + json.dumps(jsonable_encoder(EndEvent(failed=True, msg="Failed"))), + ] + ) + run_detached_monitor(monitor_client, experiment_id="test-run-id") captured = capsys.readouterr() expected = [ "============= Optimization progress (Batch #0) =============\n",