diff --git a/python_tests/artifact/test_backend.py b/python_tests/artifact/test_backend.py index 14a470e3..52fe6027 100644 --- a/python_tests/artifact/test_backend.py +++ b/python_tests/artifact/test_backend.py @@ -1,3 +1,5 @@ +import base64 +import json import tempfile from unittest.mock import MagicMock @@ -177,3 +179,73 @@ def test_successful_trial_artifact_retrieval() -> None: ) assert status == 200 assert body == b"dummy_content" + + +DUMMY_DATA_URL = ( + f"data:text/plain; charset=utf-8,{base64.b64encode(b'dummy_content').decode('utf-8')}" +) + + +def test_upload_artifact_invalid_no_trial() -> None: + storage = optuna.storages.InMemoryStorage() + + with tempfile.TemporaryDirectory() as tmpdir: + artifact_store = FileSystemArtifactStore(tmpdir) + + app = create_app(storage, artifact_store) + study = optuna.create_study(storage=storage) + + status, _, body = send_request( + app, + f"/api/artifacts/{study._study_id}/0", + "POST", + body=json.dumps({"file": DUMMY_DATA_URL}), + content_type="application/json", + ) + assert status == 500 # TODO(contramundum53): This should return 400 + + +def test_upload_artifact_invalid_complete_trial() -> None: + storage = optuna.storages.InMemoryStorage() + + with tempfile.TemporaryDirectory() as tmpdir: + artifact_store = FileSystemArtifactStore(tmpdir) + + app = create_app(storage, artifact_store) + study = optuna.create_study(storage=storage) + + study.add_trial(optuna.create_trial(value=1.0, distributions={}, params={})) + trial = study.trials[-1] + status, _, body = send_request( + app, + f"/api/artifacts/{study._study_id}/{trial._trial_id}", + "POST", + body=json.dumps({"file": DUMMY_DATA_URL}), + content_type="application/json", + ) + assert status == 400 + + +def test_upload_artifact() -> None: + storage = optuna.storages.InMemoryStorage() + + study = optuna.create_study(storage=storage) + with tempfile.TemporaryDirectory() as tmpdir: + artifact_store = FileSystemArtifactStore(tmpdir) + + app = create_app(storage, artifact_store) + + study.add_trial(optuna.create_trial(state=optuna.trial.TrialState.RUNNING)) + trial = study.trials[-1] + status, _, body = send_request( + app, + f"/api/artifacts/{study._study_id}/{trial._trial_id}", + "POST", + body=json.dumps({"file": DUMMY_DATA_URL}), + content_type="application/json", + ) + assert status == 201 + res = json.loads(body) + with open(f"{tmpdir}/{res['artifact_id']}", "r") as f: + data = f.read() + assert data == "dummy_content"