diff --git a/optuna_dashboard/_serializer.py b/optuna_dashboard/_serializer.py index 7fae23ea..04190f59 100644 --- a/optuna_dashboard/_serializer.py +++ b/optuna_dashboard/_serializer.py @@ -15,6 +15,7 @@ from . import _note as note from ._form_widget import get_form_widgets_json from ._named_objectives import get_objective_names from .artifact._backend import list_trial_artifacts +from .preferential._study import _SYSTEM_ATTR_PREFERENTIAL_STUDY if TYPE_CHECKING: @@ -143,6 +144,7 @@ def serialize_study_detail( serialized["union_user_attrs"] = [{"key": a[0], "sortable": a[1]} for a in union_user_attrs] serialized["has_intermediate_values"] = has_intermediate_values serialized["note"] = note.get_note_from_system_attrs(system_attrs, None) + serialized["is_preferential"] = system_attrs.get(_SYSTEM_ATTR_PREFERENTIAL_STUDY, False) objective_names = get_objective_names(system_attrs) if objective_names: serialized["objective_names"] = objective_names diff --git a/python_tests/test_serializers.py b/python_tests/test_serializers.py index 4ea8c9bd..7a038b08 100644 --- a/python_tests/test_serializers.py +++ b/python_tests/test_serializers.py @@ -1,19 +1,43 @@ from __future__ import annotations -from unittest import TestCase - +import optuna from optuna_dashboard._serializer import serialize_attrs +from optuna_dashboard._serializer import serialize_study_detail +from optuna_dashboard._storage import get_study_summaries +from optuna_dashboard.preferential import create_study -class SerializeAttrsTestCase(TestCase): - def test_serialize_bytes(self) -> None: - serialized = serialize_attrs({"bytes": b"This is a bytes object."}) - self.assertEqual(serialized[0]["value"], "") +def test_serialize_bytes() -> None: + serialized = serialize_attrs({"bytes": b"This is a bytes object."}) + assert serialized[0]["value"] == "" - def test_serialize_dict(self) -> None: - serialized = serialize_attrs( - { - "key": {"foo": "bar"}, - } - ) - self.assertLessEqual(len(serialized), 1) + +def test_serialize_dict() -> None: + serialized = serialize_attrs( + { + "key": {"foo": "bar"}, + } + ) + assert len(serialized) <= 1 + + +def test_get_study_detail_is_preferential() -> None: + storage = optuna.storages.InMemoryStorage() + study = create_study(storage=storage) + study_summaries = get_study_summaries(storage) + assert len(study_summaries) == 1 + + study_summary = study_summaries[0] + study_detail = serialize_study_detail(study_summary, [], study.trials, [], [], [], False) + assert study_detail["is_preferential"] + + +def test_get_study_detail_is_not_preferential() -> None: + storage = optuna.storages.InMemoryStorage() + study = optuna.create_study(storage=storage) + study_summaries = get_study_summaries(storage) + assert len(study_summaries) == 1 + + study_summary = study_summaries[0] + study_detail = serialize_study_detail(study_summary, [], study.trials, [], [], [], False) + assert not study_detail["is_preferential"] diff --git a/python_tests/wsgi_client.py b/python_tests/wsgi_client.py index 3a6b078e..84fe8aac 100644 --- a/python_tests/wsgi_client.py +++ b/python_tests/wsgi_client.py @@ -6,12 +6,21 @@ from typing import Optional from typing import Union from bottle import Bottle +from optuna_dashboard._storage import trials_cache +from optuna_dashboard._storage import trials_cache_lock +from optuna_dashboard._storage import trials_last_fetched_at if typing.TYPE_CHECKING: from _typeshed.wsgi import WSGIEnvironment +def clear_inmemory_cache() -> None: + with trials_cache_lock: + trials_cache.clear() + trials_last_fetched_at.clear() + + def create_wsgi_env( path: str, method: str, @@ -66,6 +75,8 @@ def send_request( headers = headers or {} queries = queries or {} env = create_wsgi_env(path, method, content_type, bytes_body, queries, headers) + + clear_inmemory_cache() response_body = b"" iterable_body = app(env, start_response) for b in iterable_body: