diff --git a/docs/api.rst b/docs/api.rst index a286b2c7..39051ae3 100644 --- a/docs/api.rst +++ b/docs/api.rst @@ -40,5 +40,6 @@ Artifact optuna_dashboard.artifact.upload_artifact optuna_dashboard.artifact.file_system.FileSystemBackend optuna_dashboard.artifact.boto3.Boto3Backend + optuna_dashboard.artifact.backoff.Backoff optuna_dashboard.artifact.protocol.ArtifactBackend optuna_dashboard.artifact.exceptions.ArtifactNotFound diff --git a/optuna_dashboard/artifact/backoff.py b/optuna_dashboard/artifact/backoff.py new file mode 100644 index 00000000..9d9fcf6b --- /dev/null +++ b/optuna_dashboard/artifact/backoff.py @@ -0,0 +1,110 @@ +from __future__ import annotations + +import logging +import time +from typing import TYPE_CHECKING + +from optuna_dashboard.artifact.exceptions import ArtifactNotFound + + +_logger = logging.getLogger(__name__) + + +if TYPE_CHECKING: + from typing import BinaryIO + + from optuna_dashboard.artifact.protocol import ArtifactBackend + + +class Backoff: + """An artifact backend middleware for exponential backoff. + + Example: + .. code-block:: python + + import optuna + from optuna_dashboard.artifact import upload_artifact + from optuna_dashboard.artifact.backoff import Backoff + from optuna_dashboard.artifact.boto3 import Boto3Backend + + artifact_backend = Backoff(Boto3Backend("my-bucket")) + + def objective(trial: optuna.Trial) -> float: + ... = trial.suggest_float("x", -10, 10) + file_path = generate_example_png(...) + upload_artifact(artifact_backend, trial, file_path) + return ... + """ + + def __init__( + self, + backend: ArtifactBackend, + max_retries: int = 10, + multiplier: float = 2, + min_delay: float = 0.1, + max_delay: float = 30, + ) -> None: + # Default sleep seconds: + # 0.1, 0.2, 0.4, 0.8, 1.6, 3.2, 6.4, 12.8, 25.6, 30 + self._backend = backend + assert max_retries > 0 + assert multiplier > 0 + assert min_delay > 0 + assert max_delay > min_delay + self._max_retries = max_retries + self._multiplier = multiplier + self._min_delay = min_delay + self._max_delay = max_delay + + def _get_sleep_secs(self, n_retry: int) -> float: + return min(self._min_delay * self._multiplier**n_retry, self._max_delay) + + def open(self, artifact_id: str) -> BinaryIO: + for i in range(self._max_retries): + try: + return self._backend.open(artifact_id) + except ArtifactNotFound: + raise + except Exception as e: + if i == self._max_retries - 1: + raise + else: + _logger.error(f"Failed to open artifact={artifact_id} n_retry={i}", exc_info=e) + time.sleep(self._get_sleep_secs(i)) + assert False, "must not reach here" + + def write(self, artifact_id: str, content_body: BinaryIO) -> None: + for i in range(self._max_retries): + try: + self._backend.write(artifact_id, content_body) + break + except ArtifactNotFound: + raise + except Exception as e: + if i == self._max_retries - 1: + raise + else: + _logger.error(f"Failed to open artifact={artifact_id} n_retry={i}", exc_info=e) + content_body.seek(0) + time.sleep(self._get_sleep_secs(i)) + + def remove(self, artifact_id: str) -> None: + for i in range(self._max_retries): + try: + self._backend.remove(artifact_id) + except ArtifactNotFound: + raise + except Exception as e: + if i == self._max_retries - 1: + raise + else: + _logger.error(f"Failed to delete artifact={artifact_id}", exc_info=e) + time.sleep(self._get_sleep_secs(i)) + + +if TYPE_CHECKING: + # A mypy-runtime assertion to ensure that SCSBackend + # implements all abstract methods in ArtifactBackendProtocol. + from optuna_dashboard.artifact.file_system import FileSystemBackend + + _: ArtifactBackend = Backoff(FileSystemBackend(".")) diff --git a/python_tests/artifact/__init__.py b/python_tests/artifact/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/python_tests/artifact/stubs.py b/python_tests/artifact/stubs.py new file mode 100644 index 00000000..303b5a54 --- /dev/null +++ b/python_tests/artifact/stubs.py @@ -0,0 +1,59 @@ +from __future__ import annotations + +import copy +import io +import shutil +import threading +from typing import TYPE_CHECKING + +from optuna_dashboard.artifact.exceptions import ArtifactNotFound + + +if TYPE_CHECKING: + from typing import BinaryIO + + +class FailBackend: + def open(self, artifact_id: str) -> BinaryIO: + raise Exception("something error raised") + + def write(self, artifact_id: str, content_body: BinaryIO) -> None: + raise Exception("something error raised") + + def remove(self, artifact_id: str) -> None: + raise Exception("something error raised") + + +class InMemoryBackend: + def __init__(self) -> None: + self._data: dict[str, io.BytesIO] = {} + self._lock = threading.Lock() + + def open(self, artifact_id: str) -> BinaryIO: + with self._lock: + data = self._data.get(artifact_id) + if data is None: + raise ArtifactNotFound("not found") + return copy.deepcopy(data) + + def write(self, artifact_id: str, content_body: BinaryIO) -> None: + buf = io.BytesIO() + shutil.copyfileobj(content_body, buf) + buf.seek(0) + with self._lock: + self._data[artifact_id] = buf + + def remove(self, artifact_id: str) -> None: + with self._lock: + if artifact_id not in self._data: + raise ArtifactNotFound("not found") + del self._data[artifact_id] + + +if TYPE_CHECKING: + # A mypy-runtime assertion to ensure that SCSBackend + # implements all abstract methods in ArtifactBackendProtocol. + from optuna_dashboard.artifact.protocol import ArtifactBackend + + _fail: ArtifactBackend = FailBackend() + _inmemory: ArtifactBackend = InMemoryBackend() diff --git a/python_tests/artifact/test_backoff.py b/python_tests/artifact/test_backoff.py new file mode 100644 index 00000000..5d99c1b7 --- /dev/null +++ b/python_tests/artifact/test_backoff.py @@ -0,0 +1,35 @@ +import io +import uuid + +from optuna_dashboard.artifact.backoff import Backoff + +from .stubs import FailBackend +from .stubs import InMemoryBackend + + +def test_backoff_time() -> None: + backend = Backoff( + backend=FailBackend(), + min_delay=0.1, + multiplier=10, + max_delay=10, + ) + assert backend._get_sleep_secs(0) == 0.1 + assert backend._get_sleep_secs(1) == 1 + assert backend._get_sleep_secs(2) == 10 + + +def test_read_and_write() -> None: + artifact_id = f"test-{uuid.uuid4()}" + dummy_content = b"Hello World" + + backend = Backoff( + backend=InMemoryBackend(), + min_delay=0.1, + multiplier=10, + max_delay=10, + ) + backend.write(artifact_id, io.BytesIO(dummy_content)) + with backend.open(artifact_id) as f: + actual = f.read() + assert actual == dummy_content diff --git a/python_tests/test_boto3_artifact.py b/python_tests/artifact/test_boto3.py similarity index 100% rename from python_tests/test_boto3_artifact.py rename to python_tests/artifact/test_boto3.py diff --git a/python_tests/test_file_system_artifact.py b/python_tests/artifact/test_file_system.py similarity index 100% rename from python_tests/test_file_system_artifact.py rename to python_tests/artifact/test_file_system.py