import copy
+import fcntl
import functools
import hashlib
import inspect
import json
import math
+import os
import re
-from collections.abc import Sequence
+import stat
+from collections.abc import Iterator, Sequence
+from contextlib import contextmanager
from dataclasses import dataclass
+from datetime import datetime, timezone
from enum import IntEnum
from functools import lru_cache, singledispatch
from logging import Logger
cast,
get_args,
)
+from uuid import uuid4
import numpy as np
import optuna
# Exported for cross-module use; consumers may import despite the leading
# underscore on the instance (the class ``_OptunaNamespaces`` stays private).
_OPTUNA_NAMESPACES: Final[_OptunaNamespaces] = _OptunaNamespaces()
+_OPTUNA_BEST_PARAMS_QUARANTINE_TAG: Final[str] = "corrupt"
+_OPTUNA_BEST_PARAMS_QUARANTINE_TIE_BREAK_LIMIT: Final[int] = 8
def _optuna_best_params_path(
return base_path / f"optuna-{namespace}-best-params-{pair_to_filename(pair)}.json"
+def _warn_ambiguous_legacy_optuna_best_params(
+ base_path: Path,
+ pair: str,
+ namespace: OptunaNamespace,
+ best_params_path: Path,
+ logger: Logger | None,
+) -> None:
+ """Warn when a base-only legacy best-params file shadows the pair-safe path."""
+ if logger is None:
+ return
+ legacy_best_params_path = (
+ base_path / f"optuna-{namespace}-best-params-{pair.split('/')[0]}.json"
+ )
+ if (
+ legacy_best_params_path != best_params_path
+ and legacy_best_params_path.is_file()
+ ):
+ logger.warning(
+ f"[{pair}] Ignoring ambiguous legacy Optuna {namespace} best params "
+ f"at {legacy_best_params_path}: filename does not encode the complete "
+ f"pair identity"
+ )
+
+
_OPTUNA_LABEL_BEST_PARAMS_SCHEMA_VERSION: Final[int] = 2
"""Wire format version of pair-specific Optuna label best-params JSON files.
return params
+def _quarantine_corrupt_optuna_best_params(
+ best_params_path: Path,
+ pair: str,
+ namespace: OptunaNamespace,
+ cause: Exception,
+ logger: Logger | None,
+) -> Path | None:
+ """Atomically move corrupt persisted best params out of the live path."""
+ if not best_params_path.exists():
+ return None
+ stamp = datetime.now(timezone.utc).strftime("%Y%m%dT%H%M%S%fZ")
+ quarantine_base = (
+ f"{best_params_path.name}.{_OPTUNA_BEST_PARAMS_QUARANTINE_TAG}-{stamp}"
+ )
+ for index in range(_OPTUNA_BEST_PARAMS_QUARANTINE_TIE_BREAK_LIMIT + 1):
+ suffix = "" if index == 0 else f"-{index}"
+ quarantine_path = best_params_path.with_name(f"{quarantine_base}{suffix}")
+ if not quarantine_path.exists():
+ break
+ else:
+ raise FileExistsError(best_params_path)
+ try:
+ best_params_path.rename(quarantine_path)
+ except OSError as quarantine_error:
+ if logger is not None:
+ logger.error(
+ f"[{pair}] Optuna {namespace} best params "
+ f"{best_params_path.name} quarantine failed: {quarantine_error!r}",
+ exc_info=True,
+ )
+ raise
+ if logger is not None:
+ logger.warning(
+ f"[{pair}] Optuna {namespace} best params {best_params_path.name} "
+ f"corrupt ({cause!r}); quarantined to {quarantine_path.name}; "
+ "no persisted best params recovered"
+ )
+ return quarantine_path
+
+
+@contextmanager
+def _locked_optuna_best_params(
+ best_params_path: Path,
+ *,
+ exclusive: bool,
+) -> Iterator[None]:
+ """Serialize best-params I/O using a stable lock file."""
+ lock_path = best_params_path.parent / ".optuna-best-params.lock"
+ lock_fd = os.open(
+ lock_path,
+ (os.O_RDWR if exclusive else os.O_RDONLY)
+ | os.O_CREAT
+ | os.O_CLOEXEC
+ | os.O_NOFOLLOW,
+ 0o666,
+ )
+ try:
+ if not stat.S_ISREG(os.fstat(lock_fd).st_mode):
+ raise OSError(f"Optuna best params lock {lock_path} must be a regular file")
+ fcntl.flock(lock_fd, fcntl.LOCK_EX if exclusive else fcntl.LOCK_SH)
+ yield
+ finally:
+ os.close(lock_fd)
+
+
+def _reject_optuna_best_params_symlink(best_params_path: Path) -> None:
+ """Fail closed when the live best-params path is a symlink."""
+ if best_params_path.is_symlink():
+ raise OSError(
+ f"Optuna best params path {best_params_path} must not be a symlink"
+ )
+
+
def optuna_load_best_params(
base_path: Path,
pair: str,
expected_objective_identity: str | None = None,
) -> dict[str, Any] | None:
best_params_path = _optuna_best_params_path(base_path, pair, namespace)
- if best_params_path.is_file():
- with best_params_path.open("r", encoding="utf-8") as read_file:
- best_params = json.load(read_file)
- if namespace == _OPTUNA_NAMESPACES.label:
- return _validate_optuna_label_best_params(
- best_params,
- pair,
- logger,
- expected_selection_metadata=expected_selection_metadata,
+ if not best_params_path.parent.is_dir():
+ return None
+ malformed = False
+ with _locked_optuna_best_params(best_params_path, exclusive=False):
+ _reject_optuna_best_params_symlink(best_params_path)
+ if not best_params_path.is_file():
+ _warn_ambiguous_legacy_optuna_best_params(
+ base_path, pair, namespace, best_params_path, logger
)
- if expected_objective_identity is not None:
- if (
- not isinstance(best_params, dict)
- or best_params.get("objective_identity") != expected_objective_identity
- or not isinstance(best_params.get("params"), dict)
- ):
- if logger is not None:
- logger.warning(
- f"[{pair}] Ignoring Optuna {namespace} best params: "
- f"objective identity does not match "
- f"{expected_objective_identity!r}"
- )
+ return None
+ try:
+ with best_params_path.open("r", encoding="utf-8") as read_file:
+ best_params = json.load(read_file)
+ except (json.JSONDecodeError, UnicodeDecodeError):
+ malformed = True
+ if malformed:
+ with _locked_optuna_best_params(best_params_path, exclusive=True):
+ _reject_optuna_best_params_symlink(best_params_path)
+ if not best_params_path.is_file():
return None
- return best_params["params"]
- return best_params
- legacy_best_params_path = (
- base_path / f"optuna-{namespace}-best-params-{pair.split('/')[0]}.json"
- )
- if (
- logger is not None
- and legacy_best_params_path != best_params_path
- and legacy_best_params_path.is_file()
- ):
- logger.warning(
- f"[{pair}] Ignoring ambiguous legacy Optuna {namespace} best params "
- f"at {legacy_best_params_path}: filename does not encode the complete "
- f"pair identity"
- )
- return None
+ try:
+ with best_params_path.open("r", encoding="utf-8") as read_file:
+ best_params = json.load(read_file)
+ except (json.JSONDecodeError, UnicodeDecodeError) as decode_error:
+ quarantined = _quarantine_corrupt_optuna_best_params(
+ best_params_path,
+ pair,
+ namespace,
+ decode_error,
+ logger,
+ )
+ if quarantined is None:
+ raise
+ return None
+ if namespace == _OPTUNA_NAMESPACES.label:
+ return _validate_optuna_label_best_params(
+ best_params,
+ pair,
+ logger,
+ expected_selection_metadata=expected_selection_metadata,
+ )
+ if expected_objective_identity is not None:
+ if (
+ not isinstance(best_params, dict)
+ or best_params.get("objective_identity") != expected_objective_identity
+ or not isinstance(best_params.get("params"), dict)
+ ):
+ if logger is not None:
+ logger.warning(
+ f"[{pair}] Ignoring Optuna {namespace} best params: "
+ f"objective identity does not match "
+ f"{expected_objective_identity!r}"
+ )
+ return None
+ return best_params["params"]
+ return best_params
def optuna_save_best_params(
objective_identity: str | None = None,
) -> None:
best_params_path = _optuna_best_params_path(base_path, pair, namespace)
+ temporary_path: Path | None = None
try:
if namespace == _OPTUNA_NAMESPACES.label:
best_params: dict[str, Any] = {
}
else:
best_params = params
- with best_params_path.open("w", encoding="utf-8") as write_file:
- json.dump(best_params, write_file, indent=4)
- except Exception as e:
- logger.error(
- f"[{pair}] Optuna {namespace} failed to save best params: {e!r}",
- exc_info=True,
- )
+ with _locked_optuna_best_params(best_params_path, exclusive=True):
+ _reject_optuna_best_params_symlink(best_params_path)
+ try:
+ existing_metadata = best_params_path.stat()
+ except FileNotFoundError:
+ existing_metadata = None
+ temporary_path = best_params_path.with_name(
+ f".{best_params_path.name}.{uuid4().hex}.tmp"
+ )
+ with os.fdopen(
+ os.open(
+ temporary_path,
+ os.O_WRONLY | os.O_CREAT | os.O_EXCL,
+ 0o666,
+ ),
+ mode="w",
+ encoding="utf-8",
+ ) as write_file:
+ if existing_metadata is not None:
+ temporary_metadata = os.fstat(write_file.fileno())
+ if (
+ temporary_metadata.st_uid != existing_metadata.st_uid
+ or temporary_metadata.st_gid != existing_metadata.st_gid
+ ):
+ os.fchown(
+ write_file.fileno(),
+ existing_metadata.st_uid
+ if temporary_metadata.st_uid != existing_metadata.st_uid
+ else -1,
+ existing_metadata.st_gid
+ if temporary_metadata.st_gid != existing_metadata.st_gid
+ else -1,
+ )
+ os.fchmod(
+ write_file.fileno(),
+ stat.S_IMODE(existing_metadata.st_mode),
+ )
+ json.dump(best_params, write_file, indent=4)
+ write_file.flush()
+ os.fsync(write_file.fileno())
+ os.replace(temporary_path, best_params_path)
+ temporary_path = None
+ except BaseException as error:
+ if temporary_path is not None:
+ try:
+ temporary_path.unlink(missing_ok=True)
+ except OSError as cleanup_error:
+ logger.error(
+ f"[{pair}] Optuna {namespace} best params temporary file "
+ f"{temporary_path.name} cleanup failed: {cleanup_error!r}",
+ exc_info=True,
+ )
+ if isinstance(error, Exception):
+ logger.error(
+ f"[{pair}] Optuna {namespace} failed to save best params: {error!r}",
+ exc_info=True,
+ )
raise