diff --git a/optuna_dashboard/_storage.py b/optuna_dashboard/_storage.py index ecf9c237..803aabd4 100644 --- a/optuna_dashboard/_storage.py +++ b/optuna_dashboard/_storage.py @@ -59,13 +59,6 @@ def get_trials(storage: BaseStorage, study_id: int) -> list[FrozenTrial]: return trials -def get_trial(storage: BaseStorage, study_id: int, trial_id: int) -> FrozenTrial | None: - for trial in get_trials(storage, study_id): - if trial._trial_id == trial_id: - return trial - return None - - def get_study_summaries(storage: BaseStorage) -> list[StudySummary]: if version.parse(optuna_ver) >= version.Version("3.0.0rc0.dev"): frozen_studies = storage.get_all_studies() # type: ignore diff --git a/optuna_dashboard/artifact/_backend.py b/optuna_dashboard/artifact/_backend.py index 71ea38c7..63ae8fb8 100644 --- a/optuna_dashboard/artifact/_backend.py +++ b/optuna_dashboard/artifact/_backend.py @@ -18,7 +18,6 @@ from optuna.trial import FrozenTrial from .._bottle_util import json_api_view from .._bottle_util import parse_data_uri -from .._storage import get_trial if TYPE_CHECKING: @@ -106,7 +105,7 @@ def register_artifact_route( storage.set_study_system_attr(study_id, attr_key, json.dumps(artifact)) response.status = 201 - trial = get_trial(storage, study_id, trial_id) + trial = storage.get_trial(trial_id) if trial is None: response.status = 400 return {"reason": "Invalid study_id or trial_id"} diff --git a/python_tests/artifact/test_optuna_compatibility.py b/python_tests/artifact/test_optuna_compatibility.py new file mode 100644 index 00000000..a7bb5bf0 --- /dev/null +++ b/python_tests/artifact/test_optuna_compatibility.py @@ -0,0 +1,42 @@ +from __future__ import annotations + +import tempfile + +import optuna +from optuna.version import __version__ as optuna_ver +from optuna_dashboard.artifact._backend import list_trial_artifacts +from packaging import version +import pytest + + +@pytest.mark.skipIf( + version.parse(optuna_ver) < version.Version("3.3.0.dev"), + "Artifact is not implemented yet in Optuna", +) +def test_list_optuna_artifacts() -> None: + from optuna.artifacts import FileSystemArtifactStore + from optuna.artifacts import upload_artifact + + storage = optuna.storages.InMemoryStorage() + study = optuna.create_study(storage=storage) + dummy_content = b"dummy content" + + with tempfile.TemporaryDirectory() as tmpdir: + artifact_store = FileSystemArtifactStore(tmpdir) + trial = study.ask() + + with tempfile.NamedTemporaryFile() as f: + f.write(dummy_content) + f.flush() + upload_artifact(trial, f.name, artifact_store=artifact_store) + + study.tell(trial, 0.0) + + study_system_attrs = storage.get_study_system_attrs(study._study_id) + trial = storage.get_trial(trial._trial_id) + artifact_meta_list = list_trial_artifacts(study_system_attrs, trial) + assert len(artifact_meta_list) == 1 + + artifact_id = artifact_meta_list[0]["artifact_id"] + with artifact_store.open_reader(artifact_id) as reader: + assert reader.read() == dummy_content \ No newline at end of file