diff --git a/python_tests/test_api.py b/python_tests/test_api.py index ce78da47..8aa8fac2 100644 --- a/python_tests/test_api.py +++ b/python_tests/test_api.py @@ -220,6 +220,49 @@ class APITestCase(TestCase): assert study_detail["feedback_component_type"]["output_type"] == "artifact" assert study_detail["feedback_component_type"]["artifact_key"] == "image" + def test_save_trial_user_attrs(self) -> None: + study = optuna.create_study() + trials: list[optuna.Trial] = [] + for _ in range(2): + trial = study.ask() + trials.append(trial) + + request_body = { + "user_attrs": { + "number": 0, + }, + } + + app = create_app(study._storage) + status, _, _ = send_request( + app, + f"/api/trials/{trials[0]._trial_id}/user-attrs", + "POST", + content_type="application/json", + body=json.dumps(request_body), + ) + self.assertEqual(status, 204) + + assert study.trials[0].user_attrs == request_body["user_attrs"] + assert study.trials[1].user_attrs == {} + + + def test_save_trial_user_attrs_empty(self) -> None: + study = optuna.create_study() + trial = study.ask() + + app = create_app(study._storage) + status, _, _ = send_request( + app, + f"/api/trials/{trial._trial_id}/user-attrs", + "POST", + content_type="application/json", + body=json.dumps({}), + ) + self.assertEqual(status, 400) + assert study.trials[0].user_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()