#!/usr/bin/env python3
"""``omniopt_optuna`` - drive Optuna directly through OmniOpt's framework.

This is the *user-facing* entry point for running Optuna as a special case
of OmniOpt. Under the hood it shells out to OmniOpt with ``--model OPTUNA_*``
flags, so every Optuna feature (single-objective, multi-objective,
pruning, persistent studies, remote control via ``.optuna_runner.py
study ...``) is reachable through the OmniOpt CLI surface.

Usage
-----

    # Single-objective TPE
    omniopt_optuna --max-eval=20 --parameter 'x range -10 10 int'

    # Multi-objective NSGA-II
    omniopt_optuna --model=OPTUNA_NSGAII --max-eval=20 \
        --result-names RESULT1=min RESULT2=max \
        --parameter 'x range 0 5 float'

    # Remote control: create a study, drive it from any language
    omniopt_optuna study create --workdir runs/my_study --sampler=tpe
    omniopt_optuna study add    --workdir runs/my_study \
        --params-file p.json --values '{"RESULT": 3.5}'
    omniopt_optuna study suggest --workdir runs/my_study \
        --parameters-file spec.json

    # Bypass the OmniOpt orchestrator entirely - just get the next point
    omniopt_optuna suggest --sampler=cmaes --extra-iters=1 \
        <workdir-containing-input.json>

The CLI is intentionally thin: every flag here ultimately maps to a flag
passed to ``./omniopt`` (or to ``.optuna_runner.py`` for the study subcommands).
"""

from __future__ import annotations

import argparse
import base64
import os
import subprocess
import sys
from pathlib import Path
from typing import List, Optional

THIS_DIR = Path(__file__).resolve().parent
REPO_ROOT = THIS_DIR
OMNIOPT = REPO_ROOT / "omniopt"
OPTUNA_RUNNER = REPO_ROOT / ".optuna_runner.py"


SUPPORTED_MODELS = [
    "OPTUNA_TPE",
    "OPTUNA_CMAES",
    "OPTUNA_GP",
    "OPTUNA_RANDOM",
    "OPTUNA_GRID",
    "OPTUNA_QMC",
    "OPTUNA_BruteForce",
    "OPTUNA_NSGAII",
    "OPTUNA_NSGAIII",
    "OPTUNA_MOTPE",
]


def _b64(text: str) -> str:
    return base64.b64encode(text.encode("utf-8")).decode("ascii")


def _resolve_python() -> str:
    """Pick the python interpreter that has optuna installed.

    Preference order: the active ``VIRTUAL_ENV`` python (if any), then any
    framework venv under ``~/.omniax_venvs``, then ``sys.executable``.
    """
    venv = os.environ.get("VIRTUAL_ENV")
    if venv and (Path(venv) / "bin" / "python3").exists():
        return str(Path(venv) / "bin" / "python3")
    home = Path.home()
    for pyver in ("Python_3.13.5", "Python_3.12.10", "Python_3.11.9"):
        for arch in ("x86_64", "aarch64"):
            cand = home / ".omniax_venvs" / pyver / arch / "bin" / "python3"
            if cand.exists():
                return str(cand)
    return sys.executable


def _ensure_deps() -> None:
    """Best-effort: ensure optuna + beartype are importable in the runner."""
    for mod in ("optuna", "beartype"):
        try:
            __import__(mod)
        except ImportError:
            print(
                f"omniopt_optuna: {mod} is not importable; "
                f"install it with `pip install {mod}`.",
                file=sys.stderr,
            )


def _resolve_storage(storage: Optional[str], workdir: Optional[Path]) -> Optional[str]:
    if storage:
        return storage
    if workdir is None:
        return None
    db = Path(workdir) / "optuna_study.db"
    return f"sqlite:///{db}"


def _build_omniopt_cmd(
    *,
    model: str,
    max_eval: int,
    mem_gb: int,
    time_limit: int,
    worker_timeout: int,
    num_parallel_jobs: int,
    gpus: int,
    num_random_steps: int,
    run_program: Optional[str],
    parameters: List[str],
    result_names: List[str],
    follow: bool,
    generate_all_jobs_at_once: bool,
    optuna_sampler: Optional[str],
    optuna_pruner: str,
    optuna_n_startup_trials: int,
    optuna_multivariate: bool,
    optuna_group: bool,
    optuna_constraints: bool,
    optuna_n_ei_candidates: int,
    optuna_storage: Optional[str],
    optuna_study_name: str,
    optuna_no_load_if_exists: bool,
    optuna_extra_iters: int,
    experiment_name: str,
    extra_argv: List[str],
) -> List[str]:
    """Build the ``omniopt`` argv that this entry point will run."""
    if run_program:
        run_program_b64 = _b64(run_program)
    else:
        run_program_b64 = _b64("echo 'RESULT: 0'")

    cmd: List[str] = [
        str(OMNIOPT),
        f"--model={model}",
        f"--max_eval={max_eval}",
        f"--mem_gb={mem_gb}",
        f"--time={time_limit}",
        f"--worker_timeout={worker_timeout}",
        f"--num_parallel_jobs={num_parallel_jobs}",
        f"--gpus={gpus}",
        f"--num_random_steps={num_random_steps}",
        "--run_mode=local",
        f"--run_program={run_program_b64}",
        f"--experiment_name={experiment_name}",
        f"--optuna_pruner={optuna_pruner}",
        f"--optuna_n_startup_trials={optuna_n_startup_trials}",
        f"--optuna_n_ei_candidates={optuna_n_ei_candidates}",
        f"--optuna_study_name={optuna_study_name}",
        f"--optuna_extra_iters={optuna_extra_iters}",
        "--send_anonymized_usage_stats",
    ]

    if follow:
        cmd.append("--follow")
    if generate_all_jobs_at_once:
        cmd.append("--generate_all_jobs_at_once")
    if optuna_sampler:
        cmd.append(f"--optuna_sampler={optuna_sampler}")
    if optuna_multivariate:
        cmd.append("--optuna_multivariate")
    if optuna_group:
        cmd.append("--optuna_group")
    if optuna_constraints:
        cmd.append("--optuna_constraints")
    if optuna_no_load_if_exists:
        cmd.append("--optuna_no_load_if_exists")
    storage = _resolve_storage(optuna_storage, Path.cwd())
    if storage:
        cmd.append(f"--optuna_storage={storage}")

    for r in result_names:
        cmd.extend(["--result_names", r])
    for p in parameters:
        cmd.extend(["--parameter", *p.split()])

    cmd.extend(extra_argv)
    return cmd


def _build_parser() -> argparse.ArgumentParser:
    p = argparse.ArgumentParser(
        prog="omniopt_optuna",
        description=(
            "Drive Optuna directly through OmniOpt's framework. "
            "Wraps ./omniopt --model OPTUNA_* and exposes the same study "
            "remote-control commands as .optuna_runner.py study ..."
        ),
    )
    sub = p.add_subparsers(dest="cmd", required=False)

    run = sub.add_parser(
        "run", help=argparse.SUPPRESS,
        description="Run an Optuna optimization through OmniOpt.",
    )
    _add_run_args(run)
    sub.add_parser("suggest", help=argparse.SUPPRESS,
                   description="Get just the next point without running OmniOpt.")
    sub.add_parser(
        "study",
        help="Remote-control a persistent Optuna study.",
        description=(
            "Create / inspect / drive an Optuna study that lives on disk. "
            "Use this when you want to talk to Optuna from any language, or "
            "from a process that does not import Optuna."
        ),
    )
    return p


def _add_run_args(p: argparse.ArgumentParser) -> None:
    p.add_argument("--model", default="OPTUNA_TPE", choices=SUPPORTED_MODELS,
                   help="Optuna model to use (default: OPTUNA_TPE)")
    p.add_argument("--max-eval", type=int, default=10)
    p.add_argument("--mem-gb", type=int, default=1)
    p.add_argument("--time", type=int, default=60)
    p.add_argument("--worker-timeout", type=int, default=60)
    p.add_argument("--num-parallel-jobs", type=int, default=1)
    p.add_argument("--gpus", type=int, default=0)
    p.add_argument("--num-random-steps", type=int, default=0,
                   help="How many SOBOL steps to run before Optuna (default: 0)")
    p.add_argument("--run-program", default=None,
                   help="Shell program OmniOpt runs per trial. "
                        "Use %(name)s placeholders for parameters.")
    p.add_argument("--parameter", dest="parameters", action="append", default=[],
                   help='OmniOpt-style --parameter, e.g. "x range -10 10 int". '
                        "Can be passed multiple times.")
    p.add_argument("--result-names", nargs="+", default=["RESULT=min"],
                   help='OmniOpt-style --result_names, e.g. "RESULT1=min RESULT2=max"')
    p.add_argument("--experiment-name", default="omniopt_optuna")
    p.add_argument("--follow", action="store_true")
    p.add_argument("--generate-all-jobs-at-once", action="store_true")
    p.add_argument("--optuna-sampler", default=None)
    p.add_argument("--optuna-pruner", default="none")
    p.add_argument("--optuna-n-startup-trials", type=int, default=10)
    p.add_argument("--optuna-multivariate", action="store_true")
    p.add_argument("--optuna-group", action="store_true")
    p.add_argument("--optuna-constraints", action="store_true")
    p.add_argument("--optuna-n-ei-candidates", type=int, default=0)
    p.add_argument("--optuna-storage", default=None)
    p.add_argument("--optuna-study-name", default="omniopt_study")
    p.add_argument("--optuna-no-load-if-exists", action="store_true")
    p.add_argument("--optuna-extra-iters", type=int, default=1)
    p.add_argument("--dry-run", action="store_true",
                   help="Print the OmniOpt command without running it")
    p.add_argument("extra_argv", nargs=argparse.REMAINDER,
                   help="Extra arguments appended verbatim to the OmniOpt call.")


def _run_omniopt(cmd: List[str], dry_run: bool) -> int:
    print("$ " + " ".join(cmd))
    if dry_run:
        return 0
    return subprocess.call(cmd, cwd=str(REPO_ROOT))


def _resolve_run_args(args: argparse.Namespace) -> dict:
    return {
        "model": args.model,
        "max_eval": args.max_eval,
        "mem_gb": args.mem_gb,
        "time_limit": args.time,
        "worker_timeout": args.worker_timeout,
        "num_parallel_jobs": args.num_parallel_jobs,
        "gpus": args.gpus,
        "num_random_steps": args.num_random_steps,
        "run_program": args.run_program,
        "parameters": args.parameters,
        "result_names": args.result_names,
        "follow": args.follow,
        "generate_all_jobs_at_once": args.generate_all_jobs_at_once,
        "optuna_sampler": args.optuna_sampler,
        "optuna_pruner": args.optuna_pruner,
        "optuna_n_startup_trials": args.optuna_n_startup_trials,
        "optuna_multivariate": args.optuna_multivariate,
        "optuna_group": args.optuna_group,
        "optuna_constraints": args.optuna_constraints,
        "optuna_n_ei_candidates": args.optuna_n_ei_candidates,
        "optuna_storage": args.optuna_storage,
        "optuna_study_name": args.optuna_study_name,
        "optuna_no_load_if_exists": args.optuna_no_load_if_exists,
        "optuna_extra_iters": args.optuna_extra_iters,
        "experiment_name": args.experiment_name,
        "extra_argv": list(args.extra_argv or []),
    }


def _dispatch_suggest() -> int:
    """``omniopt_optuna suggest`` -> ``.optuna_runner.py suggest``."""
    if not OPTUNA_RUNNER.exists():
        print(f"Cannot find {OPTUNA_RUNNER}", file=sys.stderr)
        return 2
    argv = sys.argv[2:]
    return subprocess.call([_resolve_python(), str(OPTUNA_RUNNER), "suggest", *argv])


def _dispatch_study() -> int:
    """``omniopt_optuna study ...`` -> ``.optuna_runner.py study ...``."""
    if not OPTUNA_RUNNER.exists():
        print(f"Cannot find {OPTUNA_RUNNER}", file=sys.stderr)
        return 2
    argv = sys.argv[2:]
    return subprocess.call([_resolve_python(), str(OPTUNA_RUNNER), "study", *argv])


def _dispatch_run(args: argparse.Namespace) -> int:
    if not OMNIOPT.exists():
        print(f"Cannot find {OMNIOPT}", file=sys.stderr)
        return 2
    if not args.parameters:
        print(
            "ERROR: at least one --parameter is required.\n"
            "Example: --parameter 'x range -10 10 int'",
            file=sys.stderr,
        )
        return 2

    cmd = _build_omniopt_cmd(**_resolve_run_args(args))
    return _run_omniopt(cmd, args.dry_run)


def main(argv: Optional[List[str]] = None) -> int:
    argv = list(argv if argv is not None else sys.argv[1:])
    if not argv or argv[0] in ("-h", "--help"):
        # Default: print top-level help
        _build_parser().print_help()
        return 0

    _ensure_deps()

    if argv[0] == "study":
        return _dispatch_study()
    if argv[0] == "suggest":
        return _dispatch_suggest()

    # Default: ``omniopt_optuna --model OPTUNA_TPE --parameter ...`` style run.
    # Build a parser that doesn't require the leading ``run`` keyword.
    parser = argparse.ArgumentParser(
        prog="omniopt_optuna",
        description=(
            "Drive Optuna directly through OmniOpt's framework. "
            "For the ``study ...`` and ``suggest ...`` subcommands see "
            "``omniopt_optuna study --help``."
        ),
    )
    _add_run_args(parser)
    args = parser.parse_args(argv)
    return _dispatch_run(args)


if __name__ == "__main__":
    try:
        sys.exit(main())
    except KeyboardInterrupt:
        sys.exit(130)
