import argparse
import copy
import itertools
import json
import sys
from pathlib import Path
from typing import Any, Dict, List, Sequence, Tuple

from simulator import run_closed_loop


def _parse_json_value(s: str) -> Any:
    try:
        return json.loads(s)
    except json.JSONDecodeError:
        return s


def _set_by_path(cfg: Dict[str, Any], path: str, value: Any) -> None:
    parts = [p for p in str(path).split(".") if p]
    if not parts:
        raise ValueError("Empty path")

    cur: Any = cfg
    for p in parts[:-1]:
        if not isinstance(cur, dict):
            raise ValueError(f"Cannot set '{path}': parent is not an object")
        if p not in cur or not isinstance(cur[p], dict):
            cur[p] = {}
        cur = cur[p]

    if not isinstance(cur, dict):
        raise ValueError(f"Cannot set '{path}': parent is not an object")
    cur[parts[-1]] = value


def _del_by_path(cfg: Dict[str, Any], path: str) -> None:
    parts = [p for p in str(path).split(".") if p]
    if not parts:
        return

    cur: Any = cfg
    for p in parts[:-1]:
        if not isinstance(cur, dict) or p not in cur:
            return
        cur = cur[p]

    if isinstance(cur, dict):
        cur.pop(parts[-1], None)


def _sanitize(s: str) -> str:
    out: List[str] = []
    for ch in str(s):
        if ch.isalnum() or ch in ("-", "_", ".", "="):
            out.append(ch)
        else:
            out.append("_")
    return "".join(out)


def _token_for(path: str, value: Any) -> str:
    vtxt = json.dumps(value, sort_keys=True, separators=(",", ":"))
    return f"{str(path).replace('.', '_')}={vtxt}"


def _grid_product(grid: Sequence[Tuple[str, Sequence[Any]]]) -> List[Dict[str, Any]]:
    if not grid:
        return [{}]

    keys = [k for (k, _) in grid]
    values = [list(vs) for (_, vs) in grid]
    out: List[Dict[str, Any]] = []
    for combo in itertools.product(*values):
        out.append({str(k): v for (k, v) in zip(keys, combo)})
    return out


def main() -> None:
    p = argparse.ArgumentParser()
    p.add_argument("--base_config", required=True)
    p.add_argument("--out_root", required=True)
    p.add_argument("--grid", action="append", nargs=2, metavar=("KEY", "JSON_LIST"), default=[])
    p.add_argument("--set", dest="sets", action="append", nargs=2, metavar=("KEY", "JSON_VALUE"), default=[])
    p.add_argument("--del", dest="dels", action="append", metavar="KEY", default=[])
    p.add_argument("--name_prefix", default=None)
    p.add_argument("--seed_base", type=int, default=None)
    p.add_argument("--seed_stride", type=int, default=1)
    p.add_argument("--dry_run", action="store_true")
    p.add_argument("--overwrite", action="store_true")
    p.add_argument("--continue_on_error", action="store_true")
    args = p.parse_args()

    base_path = Path(args.base_config).resolve()
    base_cfg = json.loads(base_path.read_text())
    if not isinstance(base_cfg, dict):
        raise ValueError("base config must be a JSON object")

    out_root = Path(args.out_root).resolve()
    if not bool(args.dry_run):
        out_root.mkdir(parents=True, exist_ok=True)

    grid: List[Tuple[str, Sequence[Any]]] = []
    for k, v in args.grid:
        vv = _parse_json_value(v)
        if not isinstance(vv, list):
            raise ValueError(f"--grid {k} expects JSON array")
        grid.append((str(k), list(vv)))

    set_overrides: List[Tuple[str, Any]] = []
    for k, v in args.sets:
        set_overrides.append((str(k), _parse_json_value(v)))

    del_paths = [str(x) for x in args.dels]

    base_seed = int(base_cfg.get("seed", 0))
    seed_base = int(args.seed_base) if args.seed_base is not None else base_seed
    seed_stride = int(args.seed_stride)
    if seed_stride < 0:
        raise ValueError("--seed_stride must be >= 0")

    prefix = str(args.name_prefix) if args.name_prefix is not None else base_path.stem

    combos = _grid_product(grid)
    total = len(combos)

    for idx, combo in enumerate(combos):
        cfg = copy.deepcopy(base_cfg)

        for k, v in set_overrides:
            _set_by_path(cfg, k, v)
        for k, v in combo.items():
            _set_by_path(cfg, k, v)
        for k in del_paths:
            _del_by_path(cfg, k)

        cfg["seed"] = int(seed_base + idx * seed_stride)

        tokens: List[str] = []
        if combo:
            for k in sorted(combo.keys()):
                tokens.append(_sanitize(_token_for(k, combo[k])))
        else:
            for k, v in set_overrides:
                tokens.append(_sanitize(_token_for(k, v)))
            for k in del_paths:
                tokens.append(_sanitize(f"del={k.replace('.', '_')}"))

        name = prefix
        if tokens:
            name = name + "__" + "__".join(tokens)
        name = name + f"__i{idx:03d}"

        out_dir = out_root / name
        if bool(args.dry_run):
            print(f"DRY_RUN {idx + 1}/{total}: {name}")
            continue

        if out_dir.exists() and not bool(args.overwrite):
            try:
                if any(out_dir.iterdir()):
                    raise FileExistsError(str(out_dir))
            except FileNotFoundError:
                pass
        out_dir.mkdir(parents=True, exist_ok=True)

        cfg["out_dir"] = str(out_dir.resolve())

        config_used_path = out_dir / "config_used.json"
        config_used_path.write_text(json.dumps(cfg, indent=2, sort_keys=True))

        print(f"RUN {idx + 1}/{total}: {name}")
        try:
            run_closed_loop(cfg, config_path=config_used_path)
        except Exception as e:
            print(f"ERROR {name}: {e}", file=sys.stderr)
            if not bool(args.continue_on_error):
                raise

    print("OK")


if __name__ == "__main__":
    main()
