diff --git a/python_tests/test_api.py b/python_tests/test_api.py index 721bdb56..6c618816 100644 --- a/python_tests/test_api.py +++ b/python_tests/test_api.py @@ -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()