From f2f5e3665728465d89fafbbd4b34074ecb75490f Mon Sep 17 00:00:00 2001 From: nabenabe0928 Date: Wed, 22 Nov 2023 06:40:15 +0100 Subject: [PATCH 1/9] Add tests for save_trial_note --- python_tests/test_api.py | 56 ++++++++++++++++++++++++++++++++++++++++ 1 file changed, 56 insertions(+) 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() From f0f876596a1ad1e98d8016a0d03732a1566c3c98 Mon Sep 17 00:00:00 2001 From: nabenabe0928 Date: Wed, 22 Nov 2023 06:47:33 +0100 Subject: [PATCH 2/9] Apply flake8 --- python_tests/test_api.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/python_tests/test_api.py b/python_tests/test_api.py index 6c618816..1bae1872 100644 --- a/python_tests/test_api.py +++ b/python_tests/test_api.py @@ -294,7 +294,7 @@ class APITestCase(TestCase): # 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"], + f"{note_str_key_prefix(0)}{0}": request_body["body"], } def test_save_trial_note(self) -> None: @@ -303,7 +303,7 @@ class APITestCase(TestCase): 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"], + f"{note_str_key_prefix(0)}{0}": request_body["body"], } def test_save_trial_note_with_wrong_version(self) -> None: From dcbd7a3c982694cfc01bcba9c305026675203173 Mon Sep 17 00:00:00 2001 From: nabenabe0928 Date: Wed, 22 Nov 2023 06:52:06 +0100 Subject: [PATCH 3/9] Apply formatter --- python_tests/test_api.py | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/python_tests/test_api.py b/python_tests/test_api.py index 1bae1872..e62ebba0 100644 --- a/python_tests/test_api.py +++ b/python_tests/test_api.py @@ -9,7 +9,8 @@ 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._note import note_str_key_prefix +from optuna_dashboard._note import 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 From 685e2846135eccfea41fc73b867ca94a961479c5 Mon Sep 17 00:00:00 2001 From: nabenabe0928 Date: Wed, 22 Nov 2023 06:54:55 +0100 Subject: [PATCH 4/9] Apply mypy --- python_tests/test_api.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/python_tests/test_api.py b/python_tests/test_api.py index e62ebba0..0250ad66 100644 --- a/python_tests/test_api.py +++ b/python_tests/test_api.py @@ -299,7 +299,7 @@ class APITestCase(TestCase): } def test_save_trial_note(self) -> None: - request_body = {"body": "Test note.", "version": 1} + request_body: dict[str, int | str] = {"body": "Test note.", "version": 1} status, study = self._test_save_trial_note(request_body) self.assertEqual(status, 204) assert study.system_attrs == { @@ -308,7 +308,7 @@ class APITestCase(TestCase): } def test_save_trial_note_with_wrong_version(self) -> None: - request_body = {"body": "Test note.", "version": 0} + request_body: dict[str, int | str] = {"body": "Test note.", "version": 0} status, study = self._test_save_trial_note(request_body) self.assertEqual(status, 409) assert study.system_attrs == {} From f29a6e9c0c57ed6c52e113dbc04d6f3aac3efff5 Mon Sep 17 00:00:00 2001 From: nabenabe0928 Date: Mon, 27 Nov 2023 07:30:37 +0100 Subject: [PATCH 5/9] Address the comments by c-bata --- python_tests/test_api.py | 9 ++++++--- 1 file changed, 6 insertions(+), 3 deletions(-) diff --git a/python_tests/test_api.py b/python_tests/test_api.py index 0250ad66..601362a1 100644 --- a/python_tests/test_api.py +++ b/python_tests/test_api.py @@ -302,21 +302,24 @@ class APITestCase(TestCase): request_body: dict[str, int | str] = {"body": "Test note.", "version": 1} status, study = self._test_save_trial_note(request_body) self.assertEqual(status, 204) - assert study.system_attrs == { + expected_system_attrs = { note_ver_key(0): request_body["version"], f"{note_str_key_prefix(0)}{0}": request_body["body"], } + for k, v in expected_system_attrs.items(): + assert k in study.system_attrs + assert study.system_attrs[k] == v def test_save_trial_note_with_wrong_version(self) -> None: request_body: dict[str, int | str] = {"body": "Test note.", "version": 0} status, study = self._test_save_trial_note(request_body) self.assertEqual(status, 409) - assert study.system_attrs == {} + assert note_ver_key(0) not in 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 == {} + assert note_ver_key(0) not in study.system_attrs @pytest.mark.skipif(sys.version_info < (3, 8), reason="BoTorch dropped Python3.7 support") def test_skip_trial(self) -> None: From 6838dd27da9d3fffe0b944952f900b60e832f879 Mon Sep 17 00:00:00 2001 From: nabenabe0928 Date: Mon, 27 Nov 2023 07:31:50 +0100 Subject: [PATCH 6/9] Rename util method in test_save_trial_note --- python_tests/test_api.py | 8 ++++---- 1 file changed, 4 insertions(+), 4 deletions(-) diff --git a/python_tests/test_api.py b/python_tests/test_api.py index 601362a1..01f4d800 100644 --- a/python_tests/test_api.py +++ b/python_tests/test_api.py @@ -263,7 +263,7 @@ class APITestCase(TestCase): self.assertEqual(status, 400) assert study.trials[0].user_attrs == {} - def _test_save_trial_note( + def _save_trial_note( self, request_body: dict[str, int | str] ) -> tuple[int, optuna.Study]: study = optuna.create_study() @@ -300,7 +300,7 @@ class APITestCase(TestCase): def test_save_trial_note(self) -> None: request_body: dict[str, int | str] = {"body": "Test note.", "version": 1} - status, study = self._test_save_trial_note(request_body) + status, study = self._save_trial_note(request_body) self.assertEqual(status, 204) expected_system_attrs = { note_ver_key(0): request_body["version"], @@ -312,12 +312,12 @@ class APITestCase(TestCase): def test_save_trial_note_with_wrong_version(self) -> None: request_body: dict[str, int | str] = {"body": "Test note.", "version": 0} - status, study = self._test_save_trial_note(request_body) + status, study = self._save_trial_note(request_body) self.assertEqual(status, 409) assert note_ver_key(0) not in study.system_attrs def test_save_trial_note_empty(self) -> None: - status, study = self._test_save_trial_note(request_body={}) + status, study = self._save_trial_note(request_body={}) self.assertEqual(status, 400) assert note_ver_key(0) not in study.system_attrs From ca98208cca3caebbf52b04c04fefa60efa2520eb Mon Sep 17 00:00:00 2001 From: nabenabe0928 Date: Mon, 27 Nov 2023 07:33:38 +0100 Subject: [PATCH 7/9] Replace assertEqual with == because we use pytest as a runner --- python_tests/test_api.py | 8 ++++---- 1 file changed, 4 insertions(+), 4 deletions(-) diff --git a/python_tests/test_api.py b/python_tests/test_api.py index 01f4d800..3c5981ea 100644 --- a/python_tests/test_api.py +++ b/python_tests/test_api.py @@ -291,7 +291,7 @@ class APITestCase(TestCase): content_type="application/json", body=json.dumps(request_body), ) - self.assertEqual(status, 204) + assert status == 204 # Check if the version 1 is deleted. assert study.system_attrs == { note_ver_key(0): request_body["version"], @@ -301,7 +301,7 @@ class APITestCase(TestCase): def test_save_trial_note(self) -> None: request_body: dict[str, int | str] = {"body": "Test note.", "version": 1} status, study = self._save_trial_note(request_body) - self.assertEqual(status, 204) + assert status == 204 expected_system_attrs = { note_ver_key(0): request_body["version"], f"{note_str_key_prefix(0)}{0}": request_body["body"], @@ -313,12 +313,12 @@ class APITestCase(TestCase): def test_save_trial_note_with_wrong_version(self) -> None: request_body: dict[str, int | str] = {"body": "Test note.", "version": 0} status, study = self._save_trial_note(request_body) - self.assertEqual(status, 409) + assert status == 409 assert note_ver_key(0) not in study.system_attrs def test_save_trial_note_empty(self) -> None: status, study = self._save_trial_note(request_body={}) - self.assertEqual(status, 400) + assert status == 400 assert note_ver_key(0) not in study.system_attrs @pytest.mark.skipif(sys.version_info < (3, 8), reason="BoTorch dropped Python3.7 support") From 616832c98b3127329ecaa29a289b10420910648a Mon Sep 17 00:00:00 2001 From: nabenabe0928 Date: Mon, 27 Nov 2023 07:35:12 +0100 Subject: [PATCH 8/9] Apply black to test_api --- python_tests/test_api.py | 4 +--- 1 file changed, 1 insertion(+), 3 deletions(-) diff --git a/python_tests/test_api.py b/python_tests/test_api.py index 3c5981ea..2647ed1f 100644 --- a/python_tests/test_api.py +++ b/python_tests/test_api.py @@ -263,9 +263,7 @@ class APITestCase(TestCase): self.assertEqual(status, 400) assert study.trials[0].user_attrs == {} - def _save_trial_note( - self, request_body: dict[str, int | str] - ) -> tuple[int, optuna.Study]: + def _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) From 09c95197d8c6c9a5c4b25dc41bd4f2e85aa749f7 Mon Sep 17 00:00:00 2001 From: nabenabe0928 Date: Mon, 27 Nov 2023 07:41:37 +0100 Subject: [PATCH 9/9] Make test_save_trial_note_overwrite more robust --- python_tests/test_api.py | 17 ++++++++++++----- 1 file changed, 12 insertions(+), 5 deletions(-) diff --git a/python_tests/test_api.py b/python_tests/test_api.py index 2647ed1f..e7210710 100644 --- a/python_tests/test_api.py +++ b/python_tests/test_api.py @@ -280,21 +280,28 @@ class APITestCase(TestCase): study = optuna.create_study() trial = study.ask() app = create_app(study._storage) + + def _get_request_body(note_version: int) -> dict[str, str | int]: + return {"body": f"Test note ver. {note_version}.", "version": note_version} + 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), + body=json.dumps(_get_request_body(note_version=ver)), ) assert 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"], + expected_request_body = _get_request_body(note_version=2) + expected_system_attrs = { + note_ver_key(trial_id=0): expected_request_body["version"], + f"{note_str_key_prefix(trial_id=0)}{0}": expected_request_body["body"], } + for k, v in expected_system_attrs.items(): + assert k in study.system_attrs + assert study.system_attrs[k] == v def test_save_trial_note(self) -> None: request_body: dict[str, int | str] = {"body": "Test note.", "version": 1}