Add tests for save_trial_note

This commit is contained in:
nabenabe0928
2023-11-22 06:40:15 +01:00
parent cd2159ea46
commit f2f5e36657
+56
View File
@@ -9,6 +9,7 @@ from optuna import get_all_study_summaries
from optuna.study import StudyDirection
from optuna_dashboard._app import create_app
from optuna_dashboard._app import create_new_study
from optuna_dashboard._note import note_str_key_prefix, note_ver_key
from optuna_dashboard._preference_setting import register_preference_feedback_component
from optuna_dashboard._preferential_history import NewHistory
from optuna_dashboard._preferential_history import remove_history
@@ -261,6 +262,61 @@ class APITestCase(TestCase):
self.assertEqual(status, 400)
assert study.trials[0].user_attrs == {}
def _test_save_trial_note(
self, request_body: dict[str, int | str]
) -> tuple[int, optuna.Study]:
study = optuna.create_study()
trial = study.ask()
app = create_app(study._storage)
status, _, _ = send_request(
app,
f"/api/studies/{study._study_id}/{trial._trial_id}/note",
"PUT",
content_type="application/json",
body=json.dumps(request_body),
)
return status, study
def test_save_trial_note_overwrite(self) -> None:
study = optuna.create_study()
trial = study.ask()
app = create_app(study._storage)
for ver in range(1, 3):
request_body = {"body": f"Test note ver. {ver}.", "version": ver}
status, _, _ = send_request(
app,
f"/api/studies/{study._study_id}/{trial._trial_id}/note",
"PUT",
content_type="application/json",
body=json.dumps(request_body),
)
self.assertEqual(status, 204)
# Check if the version 1 is deleted.
assert study.system_attrs == {
note_ver_key(0): request_body["version"],
f"{note_str_key_prefix(0)}{0}": request_body["body"],
}
def test_save_trial_note(self) -> None:
request_body = {"body": "Test note.", "version": 1}
status, study = self._test_save_trial_note(request_body)
self.assertEqual(status, 204)
assert study.system_attrs == {
note_ver_key(0): request_body["version"],
f"{note_str_key_prefix(0)}{0}": request_body["body"],
}
def test_save_trial_note_with_wrong_version(self) -> None:
request_body = {"body": "Test note.", "version": 0}
status, study = self._test_save_trial_note(request_body)
self.assertEqual(status, 409)
assert study.system_attrs == {}
def test_save_trial_note_empty(self) -> None:
status, study = self._test_save_trial_note(request_body={})
self.assertEqual(status, 400)
assert study.system_attrs == {}
@pytest.mark.skipif(sys.version_info < (3, 8), reason="BoTorch dropped Python3.7 support")
def test_skip_trial(self) -> None:
storage = optuna.storages.InMemoryStorage()