From bad4d8c3bd87e4cd25a143a32c8233e2ae4eeb54 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?H=C3=A5vard=20Berland?= Date: Thu, 27 Aug 2026 13:45:34 +0200 Subject: [PATCH 1/2] Let FastAPI perform pydantic validation of EverestConfig objects The context (knowledge of installed forward model steps etc) is injected as a dependency to FastAPI. The existing behaviour of muting all ConfigWarnings in the context-aware validation is continued in this commit. This assumes that both the client and the server is executed in an identical environment. --- src/ert/base_model_context.py | 10 ++ .../endpoints/experiment_server.py | 25 +++-- src/everest/config/everest_config.py | 11 +- src/everest/config/simulator_config.py | 7 +- src/everest/detached/client.py | 15 ++- tests/everest/test_everserver.py | 104 ++++++++++++++++-- 6 files changed, 149 insertions(+), 23 deletions(-) diff --git a/src/ert/base_model_context.py b/src/ert/base_model_context.py index 417fdeb7145..21626ef0cb7 100644 --- a/src/ert/base_model_context.py +++ b/src/ert/base_model_context.py @@ -6,6 +6,7 @@ from typing import TYPE_CHECKING, Any from pydantic import BaseModel +from pydantic_core.core_schema import ValidationInfo init_context_var = ContextVar("_init_context_var", default=None) @@ -22,6 +23,15 @@ def use_runtime_plugins(value: ErtRuntimePlugins) -> Iterator[None]: init_context_var.reset(token) +def get_runtime_plugins(info: ValidationInfo) -> ErtRuntimePlugins | None: + """Return the active runtime plugins for a pydantic validator. + + When validating through FastAPI, the context is only available + through init_context_var. + """ + return info.context or init_context_var.get() + + class BaseModelWithContextSupport(BaseModel, extra="forbid"): def __init__(__pydantic_self__, **data: Any) -> None: __pydantic_self__.__pydantic_validator__.validate_python( diff --git a/src/ert/dark_storage/endpoints/experiment_server.py b/src/ert/dark_storage/endpoints/experiment_server.py index be9e1f3b9cc..e0d573f8e29 100644 --- a/src/ert/dark_storage/endpoints/experiment_server.py +++ b/src/ert/dark_storage/endpoints/experiment_server.py @@ -10,6 +10,8 @@ import uuid import warnings from base64 import b64decode +from collections.abc import AsyncIterator +from contextlib import ExitStack from queue import SimpleQueue from typing import Annotated @@ -170,6 +172,17 @@ def verify_auth( authenticated = [Depends(verify_auth)] +async def _with_runtime_plugins() -> AsyncIterator[None]: + stack = ExitStack() + try: + stack.enter_context(warnings.catch_warnings()) + warnings.filterwarnings("ignore", category=ConfigWarning) + stack.enter_context(use_runtime_plugins(get_site_plugins())) + yield + finally: + stack.close() + + @router.get("/", dependencies=authenticated) def get_status() -> PlainTextResponse: return PlainTextResponse("EVEREST is running") @@ -198,19 +211,17 @@ def stop() -> Response: return Response("Raise STOP flag succeeded. EVEREST initiates shutdown..", 200) -@router.post("/" + EverEndpoints.START_EXPERIMENT, dependencies=authenticated) +@router.post( + "/" + EverEndpoints.START_EXPERIMENT, + dependencies=[*authenticated, Depends(_with_runtime_plugins)], +) async def start_experiment( - request: Request, + config: EverestConfig, background_tasks: BackgroundTasks, ) -> JSONResponse: experiment_id = str(uuid.uuid4()) experiment_state = ExperimentRunnerState() _experiments[experiment_id] = experiment_state - request_data = await request.json() - # Suppress already reported warnings when we re-validate with plugins - with warnings.catch_warnings(): - warnings.filterwarnings("ignore", category=ConfigWarning) - config = EverestConfig.with_plugins(request_data) runner = ExperimentRunner(config, experiment_id) try: background_tasks.add_task(runner.run) diff --git a/src/everest/config/everest_config.py b/src/everest/config/everest_config.py index b7a98ceb625..b7d62cc7908 100644 --- a/src/everest/config/everest_config.py +++ b/src/everest/config/everest_config.py @@ -32,7 +32,11 @@ from ruamel.yaml.nodes import ScalarNode from ruamel.yaml.representer import Representer -from ert.base_model_context import BaseModelWithContextSupport, use_runtime_plugins +from ert.base_model_context import ( + BaseModelWithContextSupport, + get_runtime_plugins, + use_runtime_plugins, +) from ert.config import ( ConfigWarning, EverestConstraintsConfig, @@ -574,8 +578,9 @@ def validate_forward_model_job_name_installed(self, info: ValidationInfo) -> Sel if not forward_model_jobs: return self installed_jobs_name = [job.name for job in install_jobs] - if info.context: # Add plugin jobs - installed_jobs_name += info.context.installed_forward_model_steps.keys() + runtime_plugins = get_runtime_plugins(info) + if runtime_plugins: # Add plugin jobs + installed_jobs_name += runtime_plugins.installed_forward_model_steps.keys() errors = [] for fm_job in forward_model_jobs: diff --git a/src/everest/config/simulator_config.py b/src/everest/config/simulator_config.py index 3d5d13b33a8..308bd103c4a 100644 --- a/src/everest/config/simulator_config.py +++ b/src/everest/config/simulator_config.py @@ -10,7 +10,7 @@ ) from pydantic_core.core_schema import ValidationInfo -from ert.base_model_context import BaseModelWithContextSupport +from ert.base_model_context import BaseModelWithContextSupport, get_runtime_plugins from ert.config import ConfigValidationError from ert.config.queue_config import ( LocalQueueOptions, @@ -150,8 +150,9 @@ def apply_site_or_default_queue_if_no_user_queue( queue_system = data.get("queue_system") if queue_system is None: options = None - if info.context: - options = info.context.queue_options + runtime_plugins = get_runtime_plugins(info) + if runtime_plugins: + options = runtime_plugins.queue_options defaulted_queue_options = ( options.model_dump() diff --git a/src/everest/detached/client.py b/src/everest/detached/client.py index c6a7c51c753..602781cbb52 100644 --- a/src/everest/detached/client.py +++ b/src/everest/detached/client.py @@ -120,6 +120,7 @@ def start_experiment( retries: int = 5, ) -> str: url, cert, auth = server_context + last_error: str | None = None for retry in range(retries): try: start_endpoint = f"{url}/{EverEndpoints.START_EXPERIMENT}" @@ -132,10 +133,20 @@ def start_experiment( ) response.raise_for_status() return response.json()["experiment_id"] - except Exception: + except requests.HTTPError: + last_error = response.text logger.debug(traceback.format_exc()) + if 400 <= response.status_code < 500: + break # 4xx should not trigger retries + time.sleep(retry) + except Exception: + last_error = traceback.format_exc() + logger.debug(last_error) time.sleep(retry) - raise RuntimeError("Failed to start experiment") + message = "Failed to start experiment" + if last_error: + message += f": {last_error}" + raise RuntimeError(message) def extract_errors_from_file(path: str) -> list[str]: diff --git a/tests/everest/test_everserver.py b/tests/everest/test_everserver.py index cef0490971e..de62d7d719b 100644 --- a/tests/everest/test_everserver.py +++ b/tests/everest/test_everserver.py @@ -1,5 +1,6 @@ import asyncio import logging +import warnings from base64 import b64encode from dataclasses import dataclass from pathlib import Path @@ -9,6 +10,7 @@ from unittest.mock import AsyncMock, MagicMock, Mock, patch import pytest +import yaml from fastapi.encoders import jsonable_encoder from fastapi.testclient import TestClient from starlette.websockets import WebSocketDisconnect @@ -37,6 +39,7 @@ OPT_FAILURE_REALIZATIONS, ) from everest.util._utils import get_everest_experiment +from tests.everest.utils import MIN_CONFIG, everest_config_with_defaults @pytest.fixture @@ -339,21 +342,19 @@ def test_that_multiple_started_experiments_each_receive_distinct_experiment_ids( client = TestClient(app) mock_runner = MagicMock() mock_runner.run = AsyncMock() - with ( - patch.object(EverestConfig, "with_plugins", return_value=MagicMock()), - patch( - "ert.dark_storage.endpoints.experiment_server.ExperimentRunner", - return_value=mock_runner, - ), + config_body = everest_config_with_defaults().to_dict() + with patch( + "ert.dark_storage.endpoints.experiment_server.ExperimentRunner", + return_value=mock_runner, ): r1 = client.post( "/experiment_server/start_experiment", - json={"type": "everest_config"}, + json=config_body, headers=auth_headers, ) r2 = client.post( "/experiment_server/start_experiment", - json={"type": "everest_config"}, + json=config_body, headers=auth_headers, ) assert r1.status_code == r2.status_code == 200 @@ -373,6 +374,93 @@ def test_that_multiple_started_experiments_each_receive_distinct_experiment_ids( _experiments.update(original) +def test_that_start_experiment_with_incomplete_schema_returns_422(monkeypatch): + monkeypatch.setenv("ERT_STORAGE_TOKEN", "password") + credentials = b64encode(b"username:password").decode() + auth_headers = {"Authorization": f"Basic {credentials}"} + client = TestClient(app) + + # Missing all required fields (controls, objective_functions, + # config_path, model): schema validation should reject this before + # the endpoint body (and thus ExperimentRunner) is ever reached. + response = client.post( + "/experiment_server/start_experiment", + json={}, + headers=auth_headers, + ) + + assert response.status_code == 422 + missing_fields = { + tuple(error["loc"]) + for error in response.json()["detail"] + if error["type"] == "missing" + } + assert missing_fields == { + ("body", "controls"), + ("body", "objective_functions"), + ("body", "config_path"), + ("body", "model"), + } + + +def test_that_start_experiment_with_unknown_forward_model_job_returns_422( + monkeypatch, +): + monkeypatch.setenv("ERT_STORAGE_TOKEN", "password") + credentials = b64encode(b"username:password").decode() + auth_headers = {"Authorization": f"Basic {credentials}"} + client = TestClient(app) + + config_body = yaml.safe_load(MIN_CONFIG) | { + "forward_model": ["totally_unknown_job_xyz"] + } + + response = client.post( + "/experiment_server/start_experiment", + json=config_body, + headers=auth_headers, + ) + + assert response.status_code == 422 + assert "unknown job totally_unknown_job_xyz" in response.text + + +def test_that_start_experiment_mutes_config_warnings(monkeypatch): + monkeypatch.setenv("ERT_STORAGE_TOKEN", "password") + credentials = b64encode(b"username:password").decode() + auth_headers = {"Authorization": f"Basic {credentials}"} + client = TestClient(app) + mock_runner = MagicMock() + mock_runner.run = AsyncMock() + + config_body = yaml.safe_load(MIN_CONFIG) + + def _raise_a_config_warning(*args, **kwargs): + warnings.warn("Forced test ConfigWarning", category=ConfigWarning, stacklevel=2) + + monkeypatch.setattr( + "everest.config.everest_config.validate_forward_model_configs", + _raise_a_config_warning, + ) + + with ( + patch( + "ert.dark_storage.endpoints.experiment_server.ExperimentRunner", + return_value=mock_runner, + ), + warnings.catch_warnings(record=True) as caught_warnings, + ): + warnings.simplefilter("always") + response = client.post( + "/experiment_server/start_experiment", + json=config_body, + headers=auth_headers, + ) + + assert response.status_code == 200 + assert not any(issubclass(w.category, ConfigWarning) for w in caught_warnings) + + async def test_websocket_no_events_on_connect(setup_client): events = [] client, subs, experiment_id = setup_client(events) From 88d2b2d11c2330a3f2770270a626a9ec22aa9cab Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?H=C3=A5vard=20Berland?= Date: Thu, 3 Sep 2026 09:51:30 +0200 Subject: [PATCH 2/2] FIXUP: Extract authorized_client fixture in test_everserver --- tests/everest/test_everserver.py | 36 ++++++++++++++------------------ 1 file changed, 16 insertions(+), 20 deletions(-) diff --git a/tests/everest/test_everserver.py b/tests/everest/test_everserver.py index de62d7d719b..3b42f800fb6 100644 --- a/tests/everest/test_everserver.py +++ b/tests/everest/test_everserver.py @@ -42,6 +42,14 @@ from tests.everest.utils import MIN_CONFIG, everest_config_with_defaults +@pytest.fixture +def authorized_client(monkeypatch): + monkeypatch.setenv("ERT_STORAGE_TOKEN", "password") + credentials = b64encode(b"username:password").decode() + auth_headers = {"Authorization": f"Basic {credentials}"} + return TestClient(app), auth_headers + + @pytest.fixture def setup_client(monkeypatch): original = dict(_experiments) @@ -331,15 +339,12 @@ class TestEvent: def test_that_multiple_started_experiments_each_receive_distinct_experiment_ids( - monkeypatch, + authorized_client, ): - monkeypatch.setenv("ERT_STORAGE_TOKEN", "password") + client, auth_headers = authorized_client original = dict(_experiments) _experiments.clear() try: - credentials = b64encode(b"username:password").decode() - auth_headers = {"Authorization": f"Basic {credentials}"} - client = TestClient(app) mock_runner = MagicMock() mock_runner.run = AsyncMock() config_body = everest_config_with_defaults().to_dict() @@ -374,11 +379,8 @@ def test_that_multiple_started_experiments_each_receive_distinct_experiment_ids( _experiments.update(original) -def test_that_start_experiment_with_incomplete_schema_returns_422(monkeypatch): - monkeypatch.setenv("ERT_STORAGE_TOKEN", "password") - credentials = b64encode(b"username:password").decode() - auth_headers = {"Authorization": f"Basic {credentials}"} - client = TestClient(app) +def test_that_start_experiment_with_incomplete_schema_returns_422(authorized_client): + client, auth_headers = authorized_client # Missing all required fields (controls, objective_functions, # config_path, model): schema validation should reject this before @@ -404,12 +406,9 @@ def test_that_start_experiment_with_incomplete_schema_returns_422(monkeypatch): def test_that_start_experiment_with_unknown_forward_model_job_returns_422( - monkeypatch, + authorized_client, ): - monkeypatch.setenv("ERT_STORAGE_TOKEN", "password") - credentials = b64encode(b"username:password").decode() - auth_headers = {"Authorization": f"Basic {credentials}"} - client = TestClient(app) + client, auth_headers = authorized_client config_body = yaml.safe_load(MIN_CONFIG) | { "forward_model": ["totally_unknown_job_xyz"] @@ -425,11 +424,8 @@ def test_that_start_experiment_with_unknown_forward_model_job_returns_422( assert "unknown job totally_unknown_job_xyz" in response.text -def test_that_start_experiment_mutes_config_warnings(monkeypatch): - monkeypatch.setenv("ERT_STORAGE_TOKEN", "password") - credentials = b64encode(b"username:password").decode() - auth_headers = {"Authorization": f"Basic {credentials}"} - client = TestClient(app) +def test_that_start_experiment_mutes_config_warnings(authorized_client, monkeypatch): + client, auth_headers = authorized_client mock_runner = MagicMock() mock_runner.run = AsyncMock()