From 0b087b17d0201c743f996ab60eb6d0b672c5fb9c Mon Sep 17 00:00:00 2001 From: HideakiImamura Date: Fri, 17 Nov 2023 14:22:13 +0900 Subject: [PATCH 1/2] Add unittests for tell_trial --- python_tests/test_api.py | 119 +++++++++++++++++++++++++++++++++++++++ 1 file changed, 119 insertions(+) diff --git a/python_tests/test_api.py b/python_tests/test_api.py index bde7838a..eb96adda 100644 --- a/python_tests/test_api.py +++ b/python_tests/test_api.py @@ -399,6 +399,125 @@ class APITestCase(TestCase): ) self.assertEqual(status, 404) + def test_tell_trial_complete(self) -> None: + storage = optuna.storages.InMemoryStorage() + study = optuna.create_study(storage=storage) + trial_id = study.ask()._trial_id + + app = create_app(storage) + status, _, _ = send_request( + app, + f"/api/trials/{trial_id}/tell", + "POST", + body=json.dumps( + { + "state": "Complete", + "values": [0, 1, 2], + } + ), + content_type="application/json", + ) + self.assertEqual(status, 204) + trial = storage.get_trial(trial_id) + assert trial.state == optuna.trial.TrialState.COMPLETE + assert trial.values == [0, 1, 2] + + def test_tell_trial_fail(self) -> None: + storage = optuna.storages.InMemoryStorage() + study = optuna.create_study(storage=storage) + trial_id = study.ask()._trial_id + + app = create_app(storage) + status, _, _ = send_request( + app, + f"/api/trials/{trial_id}/tell", + "POST", + body=json.dumps( + { + "state": "Fail", + } + ), + content_type="application/json", + ) + self.assertEqual(status, 204) + trial = storage.get_trial(trial_id) + assert trial.state == optuna.trial.TrialState.FAIL + + def test_tell_trial_with_no_state(self) -> None: + storage = optuna.storages.InMemoryStorage() + study = optuna.create_study(storage=storage) + trial_id = study.ask()._trial_id + + app = create_app(storage) + status, _, _ = send_request( + app, + f"/api/trials/{trial_id}/tell", + "POST", + body=json.dumps({}), + content_type="application/json", + ) + self.assertEqual(status, 400) + + def test_tell_trial_with_invalid_state(self) -> None: + storage = optuna.storages.InMemoryStorage() + study = optuna.create_study(storage=storage) + for state in ["Pruned", "Running", "Waiting", "Invalid"]: + trial_id = study.ask()._trial_id + + app = create_app(storage) + status, _, _ = send_request( + app, + f"/api/trials/{trial_id}/tell", + "POST", + body=json.dumps( + { + "state": state, + } + ), + content_type="application/json", + ) + self.assertEqual(status, 400) + + def test_tell_trial_with_no_values(self) -> None: + storage = optuna.storages.InMemoryStorage() + study = optuna.create_study(storage=storage) + trial_id = study.ask()._trial_id + + app = create_app(storage) + status, _, _ = send_request( + app, + f"/api/trials/{trial_id}/tell", + "POST", + body=json.dumps( + { + "state": "Complete", + } + ), + content_type="application/json", + ) + self.assertEqual(status, 400) + + def test_tell_trial_with_invalid_values(self) -> None: + storage = optuna.storages.InMemoryStorage() + study = optuna.create_study(storage=storage) + for values in [1.0, ["foo"]]: + trial_id = study.ask()._trial_id + + app = create_app(storage) + status, _, _ = send_request( + app, + f"/api/trials/{trial_id}/tell", + "POST", + body=json.dumps( + { + "state": "Complete", + "values": values, + } + ), + content_type="application/json", + ) + self.assertEqual(status, 400) + class BottleRequestHookTestCase(TestCase): def test_ignore_trailing_slashes(self) -> None: From 6d0fef3ff279cc239e8349158535eff1015f69a2 Mon Sep 17 00:00:00 2001 From: HideakiImamura Date: Fri, 17 Nov 2023 16:42:39 +0900 Subject: [PATCH 2/2] Use subtest --- python_tests/test_api.py | 54 ++++++++++++++++++++-------------------- 1 file changed, 27 insertions(+), 27 deletions(-) diff --git a/python_tests/test_api.py b/python_tests/test_api.py index eb96adda..ce78da47 100644 --- a/python_tests/test_api.py +++ b/python_tests/test_api.py @@ -463,20 +463,20 @@ class APITestCase(TestCase): study = optuna.create_study(storage=storage) for state in ["Pruned", "Running", "Waiting", "Invalid"]: trial_id = study.ask()._trial_id - app = create_app(storage) - status, _, _ = send_request( - app, - f"/api/trials/{trial_id}/tell", - "POST", - body=json.dumps( - { - "state": state, - } - ), - content_type="application/json", - ) - self.assertEqual(status, 400) + with self.subTest(state=state): + status, _, _ = send_request( + app, + f"/api/trials/{trial_id}/tell", + "POST", + body=json.dumps( + { + "state": state, + } + ), + content_type="application/json", + ) + self.assertEqual(status, 400) def test_tell_trial_with_no_values(self) -> None: storage = optuna.storages.InMemoryStorage() @@ -502,21 +502,21 @@ class APITestCase(TestCase): study = optuna.create_study(storage=storage) for values in [1.0, ["foo"]]: trial_id = study.ask()._trial_id - app = create_app(storage) - status, _, _ = send_request( - app, - f"/api/trials/{trial_id}/tell", - "POST", - body=json.dumps( - { - "state": "Complete", - "values": values, - } - ), - content_type="application/json", - ) - self.assertEqual(status, 400) + with self.subTest(values=values): + status, _, _ = send_request( + app, + f"/api/trials/{trial_id}/tell", + "POST", + body=json.dumps( + { + "state": "Complete", + "values": values, + } + ), + content_type="application/json", + ) + self.assertEqual(status, 400) class BottleRequestHookTestCase(TestCase):