import argparse
import json
import math
import os
import platform
import random
import re
import shutil
import signal
import subprocess
import time
from dataclasses import dataclass
from datetime import datetime, timezone
from pathlib import Path

from score_run import decide, load_reference, parse_run_log


ROOT = Path(__file__).resolve().parent
if os.name == "nt":
    CODEX_CMD = Path.home() / "AppData" / "Roaming" / "npm" / "codex.cmd"
else:
    CODEX_CMD = Path("/usr/bin/codex")
PROXY_BEST_PATH = ROOT / "proxy_best_run.json"
TRUTH_REFERENCE_PATH = ROOT / "current_best_run.json"
TRUTH_OVERRIDE_BEST_PATH = ROOT / "truth_override_best.json"
STATE_PATH = ROOT / "proxy_search_state.json"
PROXY_RESULTS_PATH = ROOT / "proxy_results.tsv"
TRUTH_RESULTS_PATH = ROOT / "truth_promotions.tsv"
PROXY_RUNS_DIR = ROOT / "proxy_runs"
TRUTH_RUNS_DIR = ROOT / "truth_runs"
CHECKPOINT_PATH = ROOT / "checkpoint_pre_eval.pt"
CODEX_CANDIDATE_SCHEMA_PATH = ROOT / "codex_candidate_schema.json"
CANDIDATE_WORKTREES_DIR = ROOT / "candidate_worktrees"
ALLOWED_OVERRIDE_KEYS = [
    "AUTORESEARCH_ADAMW_VARIANT_LM_HEAD",
    "AUTORESEARCH_ADAMW_VARIANT_WTE",
    "AUTORESEARCH_ADAMW_VARIANT_VALUE_EMBEDS",
    "AUTORESEARCH_ADAMW_VARIANT_RESID",
    "AUTORESEARCH_ADAMW_VARIANT_X0",
    "AUTORESEARCH_VALUE_EMBEDS_BETAS",
    "AUTORESEARCH_VALUE_EMBEDS_EPS",
    "AUTORESEARCH_MUON_SALIENCE_W_N",
    "AUTORESEARCH_MUON_SALIENCE_W_A",
    "AUTORESEARCH_MUON_SALIENCE_W_F",
    "AUTORESEARCH_MUON_SALIENCE_GATE_MAX",
    "AUTORESEARCH_SALIENCEW_PROMOTIVE_GAIN",
    "AUTORESEARCH_SALIENCEW_AVERSIVE_GAIN",
    "AUTORESEARCH_SALIENCEW_CONFLICT_GAIN",
    "AUTORESEARCH_SALIENCEW_GATE_MIN",
    "AUTORESEARCH_SALIENCEW_GATE_MAX",
    "AUTORESEARCH_MATRIX_LR",
    "AUTORESEARCH_WARMDOWN_RATIO",
    "AUTORESEARCH_FINAL_LR_FRAC",
]


@dataclass
class Candidate:
    family: str
    part: str
    cls: str
    note: str
    overrides: dict[str, str]
    proposal_mode: str = "overrides"
    edits: list[dict[str, str]] | None = None


FAMILY_PRIORS = {
    "wte_variant": 1.00,
    "lm_head_variant": 0.95,
    "value_embeds_variant": 0.95,
    "scalar_variants": 0.90,
    "value_embeds_betas": 0.85,
    "value_embeds_eps": 0.80,
    "muon_gate": 1.10,
    "saliencew_gains": 1.05,
    "schedule": 0.75,
}


def now_utc():
    return datetime.now(timezone.utc).strftime("%Y-%m-%dT%H:%M:%SZ")


def ensure_tsv(path: Path, header: str):
    if not path.exists():
        path.write_text(header + "\n", encoding="utf-8")


def build_train_command(cwd: Path, regime: str):
    venv_dirs = [
        ROOT / ".venv",
        cwd / ".venv",
    ]
    python_names = ["python.exe"] if os.name == "nt" else ["python", "python3"]
    for venv_dir in venv_dirs:
        for python_name in python_names:
            candidate = venv_dir / ("Scripts" if os.name == "nt" else "bin") / python_name
            if candidate.exists():
                return [str(candidate), "train.py", "--regime", regime]
    return ["uv", "run", "train.py", "--regime", regime]


def run_with_tree_timeout(cmd, cwd: Path, timeout_seconds: float | None, stdout=None, stderr=None, input_text: str | None = None):
    proc = subprocess.Popen(
        cmd,
        cwd=cwd,
        stdout=stdout,
        stderr=stderr,
        stdin=subprocess.PIPE if input_text is not None else None,
        text=True,
        start_new_session=(os.name != "nt"),
    )
    try:
        if input_text is not None and proc.stdin is not None:
            proc.stdin.write(input_text)
            proc.stdin.close()
        proc.wait(timeout=timeout_seconds)
        return proc.returncode
    except subprocess.TimeoutExpired:
        if os.name == "nt":
            subprocess.run(["taskkill", "/F", "/T", "/PID", str(proc.pid)], stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL)
        else:
            try:
                os.killpg(proc.pid, signal.SIGKILL)
            except ProcessLookupError:
                pass
        return 124


def load_state():
    if not STATE_PATH.exists():
        return {
            "iteration": 0,
            "proxy_keeps": 0,
            "family_stats": {name: {"trials": 0, "keeps": 0} for name in FAMILY_PRIORS},
        }
    return json.loads(STATE_PATH.read_text(encoding="utf-8"))


def save_state(state):
    STATE_PATH.write_text(json.dumps(state, indent=2, sort_keys=True), encoding="utf-8")


def get_override(overrides, key, default):
    return overrides.get(key, default)


def choose_family(state):
    total_trials = sum(item["trials"] for item in state["family_stats"].values())
    best_name = None
    best_score = None
    for name, prior in FAMILY_PRIORS.items():
        stats = state["family_stats"].setdefault(name, {"trials": 0, "keeps": 0})
        trials = stats["trials"]
        keeps = stats["keeps"]
        win_rate = keeps / trials if trials else 0.0
        explore = math.sqrt(math.log(total_trials + 2) / (trials + 1))
        score = prior + 0.75 * win_rate + 0.60 * explore
        if best_score is None or score > best_score:
            best_name = name
            best_score = score
    return best_name


def load_recent_proxy_rows(limit=8):
    if not PROXY_RESULTS_PATH.exists():
        return []
    lines = PROXY_RESULTS_PATH.read_text(encoding="utf-8").splitlines()
    if len(lines) <= 1:
        return []
    header = lines[0].split("\t")
    rows = []
    for line in lines[-limit:]:
        parts = line.split("\t")
        rows.append(dict(zip(header, parts)))
    return rows


def build_codex_prompt(best_proxy, recent_rows):
    recent_trimmed = recent_rows[-4:]
    return f"""Return exactly one JSON object and nothing else.

Do not acknowledge. Do not explain. Do not use markdown fences.

You are choosing the next candidate for a local AlphaEvolve-style overnight search loop over model/optimizer parameter overrides.

Allowed override keys:
{json.dumps(ALLOWED_OVERRIDE_KEYS, indent=2)}

Current best proxy reference:
{json.dumps(best_proxy, indent=2, sort_keys=True)}

Recent proxy results:
{json.dumps(recent_trimmed, indent=2, sort_keys=True)}

Constraints:
- mutate only 1 to 5 override keys
- make one coherent mechanism-first idea
- prefer optimizer / gating / value-embedding / schedule changes before broad architecture churn
- avoid obvious repeat losers unless there is a concrete interaction hypothesis
- all override values must be strings

Required JSON shape:
{{
  "family": "short_family_name",
  "part": "attack_matrix_part_name",
  "class": "Gate|Partition|Replace|Aux",
  "note": "one-line experiment description",
  "overrides": {{
    "ALLOWED_KEY": "string_value"
  }}
}}

Produce the candidate now.
"""


def extract_json_object(raw_text):
    stripped = raw_text.strip()
    fenced = re.search(r"```(?:json)?\s*(\{.*\})\s*```", stripped, flags=re.DOTALL)
    if fenced:
        stripped = fenced.group(1)
    start = stripped.find("{")
    end = stripped.rfind("}")
    if start == -1 or end == -1 or end <= start:
        raise ValueError("No JSON object found in Codex proposal")
    return json.loads(stripped[start : end + 1])


def normalize_codex_candidate(payload):
    overrides = payload.get("overrides", {})
    cleaned = {str(k): str(v) for k, v in overrides.items() if k in ALLOWED_OVERRIDE_KEYS}
    if not cleaned:
        raise ValueError("Codex proposal did not include any allowed overrides")
    return Candidate(
        family=str(payload["family"]),
        part=str(payload["part"]),
        cls=str(payload["class"]),
        note=str(payload["note"]),
        overrides=cleaned,
    )


def build_codex_patch_prompt(best_proxy, recent_rows):
    recent_trimmed = recent_rows[-4:]
    return f"""Return exactly one JSON object and nothing else.

You are proposing a code mutation for `train.py` inside an AlphaEvolve-style overnight search loop.

Current best proxy reference:
{json.dumps(best_proxy, indent=2, sort_keys=True)}

Recent proxy results:
{json.dumps(recent_trimmed, indent=2, sort_keys=True)}

Rules:
- modify `train.py` only
- propose 1 to 3 exact string replacements
- each `old` snippet must exist exactly in the current `train.py`
- keep the proxy/truth regime infrastructure intact
- prefer optimizer, gate, value-embedding, schedule, or OOD-sparsity edits before broad architecture churn
- make one coherent experiment, not a kitchen-sink rewrite

Required JSON shape:
{{
  "family": "short_family_name",
  "part": "attack_matrix_part_name",
  "class": "Gate|Partition|Replace|Aux",
  "note": "one-line experiment description",
  "edits": [
    {{
      "old": "exact existing snippet from train.py",
      "new": "replacement snippet"
    }}
  ]
}}

Produce the candidate now.
"""


def files_differ(path_a: Path, path_b: Path):
    return path_a.read_text(encoding="utf-8") != path_b.read_text(encoding="utf-8")


def apply_edit_instructions(base_text: str, edits):
    updated = base_text
    for edit in edits:
        old = edit["old"]
        new = edit["new"]
        if old not in updated:
            raise ValueError("Codex patch references a snippet that does not exist in train.py")
        updated = updated.replace(old, new, 1)
    return updated


def create_candidate_worktree(run_id: str):
    candidate_dir = CANDIDATE_WORKTREES_DIR / run_id
    if candidate_dir.exists():
        subprocess.run(["git", "worktree", "remove", "--force", str(candidate_dir)], cwd=ROOT, stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL)
    subprocess.run(["git", "worktree", "add", "--detach", str(candidate_dir), "HEAD"], cwd=ROOT, stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL, check=True)
    champion_train = ROOT / "train.py"
    candidate_train = candidate_dir / "train.py"
    candidate_train.write_text(champion_train.read_text(encoding="utf-8"), encoding="utf-8")
    baseline_path = PROXY_RUNS_DIR / f"{run_id}_baseline_train.py"
    baseline_path.write_text(candidate_train.read_text(encoding="utf-8"), encoding="utf-8")
    return candidate_dir, baseline_path


def cleanup_stale_candidate_worktrees():
    if not CANDIDATE_WORKTREES_DIR.exists():
        return
    for child in CANDIDATE_WORKTREES_DIR.iterdir():
        if child.is_dir():
            rc = subprocess.run(["git", "worktree", "remove", "--force", str(child)], cwd=ROOT, stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL).returncode
            if rc != 0 and child.exists():
                shutil.rmtree(child, ignore_errors=True)


def remove_candidate_worktree(candidate_dir: Path):
    rc = subprocess.run(["git", "worktree", "remove", "--force", str(candidate_dir)], cwd=ROOT, stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL).returncode
    if rc != 0 and candidate_dir.exists():
        shutil.rmtree(candidate_dir, ignore_errors=True)


def archive_candidate_patch(run_id: str, baseline_path: Path, candidate_train_path: Path):
    patch_path = PROXY_RUNS_DIR / f"{run_id}_patch.diff"
    with patch_path.open("w", encoding="utf-8") as handle:
        subprocess.run(
            ["git", "diff", "--no-index", "--", str(baseline_path), str(candidate_train_path)],
            cwd=ROOT,
            stdout=handle,
            stderr=subprocess.DEVNULL,
        )
    return patch_path


def propose_code_patch_codex(best_proxy, recent_rows, run_id, timeout_seconds: float):
    candidate_dir, baseline_path = create_candidate_worktree(run_id)
    summary_path = PROXY_RUNS_DIR / f"{run_id}_patch_summary.txt"
    patch_prompt = build_codex_patch_prompt(best_proxy, recent_rows)
    candidate_train = candidate_dir / "train.py"
    last_error = None
    try:
        for attempt in range(3):
            candidate_train.write_text(baseline_path.read_text(encoding="utf-8"), encoding="utf-8")
            if summary_path.exists():
                summary_path.unlink()
            attempt_prompt = patch_prompt
            if last_error is not None:
                attempt_prompt += "\nPrevious attempt failed: " + last_error
            cmd = [
                str(CODEX_CMD),
                "exec",
                "-C",
                str(ROOT),
                "-m",
                "gpt-5.4",
                "-c",
                'model_reasoning_effort="high"',
                "-s",
                "read-only",
                "-o",
                str(summary_path),
                "-",
            ]
            rc = run_with_tree_timeout(
                cmd,
                cwd=ROOT,
                timeout_seconds=timeout_seconds,
                stdout=subprocess.DEVNULL,
                stderr=subprocess.DEVNULL,
                input_text=attempt_prompt,
            )
            if rc == 124:
                last_error = f"Codex patch proposal timed out after {timeout_seconds}s"
                continue
            if not summary_path.exists():
                last_error = "Codex did not produce a patch proposal"
                continue
            payload = extract_json_object(summary_path.read_text(encoding="utf-8"))
            edits = payload.get("edits")
            if not edits:
                last_error = "Codex did not provide any edits"
                continue
            baseline_text = baseline_path.read_text(encoding="utf-8")
            try:
                updated_text = apply_edit_instructions(baseline_text, edits)
            except Exception as exc:
                last_error = str(exc)
                continue
            candidate_train.write_text(updated_text, encoding="utf-8")
            if not files_differ(baseline_path, candidate_train):
                last_error = "Codex edits did not change train.py"
                continue
            patch_path = archive_candidate_patch(run_id, baseline_path, candidate_train)
            summary = str(payload.get("note", "Codex edited train.py"))
            candidate = Candidate(
                family=str(payload.get("family", "codex_patch")),
                part=str(payload.get("part", "train_py_experiment")),
                cls=str(payload.get("class", "Replace")),
                note=summary,
                overrides={},
                proposal_mode="code_patch",
                edits=edits,
            )
            return candidate, candidate_dir, baseline_path, patch_path, summary_path
        raise RuntimeError(f"Codex patch proposal failed after retries: {last_error}")
    except Exception:
        remove_candidate_worktree(candidate_dir)
        raise


def propose_candidate_codex(best_proxy, recent_rows, run_id, timeout_seconds: float):
    proposal_path = PROXY_RUNS_DIR / f"{run_id}_proposal.json"
    prompt = build_codex_prompt(best_proxy, recent_rows)
    last_error = None
    for attempt in range(3):
        if proposal_path.exists():
            proposal_path.unlink()
        attempt_prompt = prompt
        if last_error is not None:
            attempt_prompt += (
                "\n\nPrevious proposal was invalid: "
                + last_error
                + "\nRetry and return valid JSON with at least one allowed override key."
            )
        cmd = [
            str(CODEX_CMD),
            "exec",
            "-C",
            str(ROOT),
            "-m",
            "gpt-5.4",
            "-c",
            'model_reasoning_effort="high"',
            "-s",
            "read-only",
            "-o",
            str(proposal_path),
            "-",
        ]
        rc = run_with_tree_timeout(
            cmd,
            cwd=ROOT,
            timeout_seconds=timeout_seconds,
            stdout=subprocess.DEVNULL,
            stderr=subprocess.DEVNULL,
            input_text=attempt_prompt,
        )
        if rc == 124:
            last_error = f"Codex proposal timed out after {timeout_seconds}s"
            continue
        if not proposal_path.exists():
            last_error = "no proposal file produced"
            continue
        try:
            payload = extract_json_object(proposal_path.read_text(encoding="utf-8"))
            return normalize_codex_candidate(payload)
        except Exception as exc:
            last_error = str(exc)
            continue
    raise RuntimeError(f"Codex proposal failed after retries: {last_error}")


def mutate_choice(rng, values, current):
    pool = [value for value in values if value != current]
    return rng.choice(pool) if pool else current


def propose_candidate(family, base_overrides, rng):
    if family == "wte_variant":
        current = get_override(base_overrides, "AUTORESEARCH_ADAMW_VARIANT_WTE", "saliencew")
        new = mutate_choice(rng, ["adamw", "saliencew"], current)
        return Candidate(family, "token_embedding_optimizer", "Partition", f"set WTE optimizer to {new}", {
            **base_overrides,
            "AUTORESEARCH_ADAMW_VARIANT_WTE": new,
        })
    if family == "lm_head_variant":
        current = get_override(base_overrides, "AUTORESEARCH_ADAMW_VARIANT_LM_HEAD", "saliencew")
        new = mutate_choice(rng, ["adamw", "saliencew"], current)
        return Candidate(family, "lm_head_optimizer", "Partition", f"set LM head optimizer to {new}", {
            **base_overrides,
            "AUTORESEARCH_ADAMW_VARIANT_LM_HEAD": new,
        })
    if family == "value_embeds_variant":
        current = get_override(base_overrides, "AUTORESEARCH_ADAMW_VARIANT_VALUE_EMBEDS", "saliencew")
        new = mutate_choice(rng, ["adamw", "saliencew"], current)
        return Candidate(family, "value_embedding_optimizer", "Partition", f"set value_embeds optimizer to {new}", {
            **base_overrides,
            "AUTORESEARCH_ADAMW_VARIANT_VALUE_EMBEDS": new,
        })
    if family == "scalar_variants":
        current_resid = get_override(base_overrides, "AUTORESEARCH_ADAMW_VARIANT_RESID", "adamw")
        current_x0 = get_override(base_overrides, "AUTORESEARCH_ADAMW_VARIANT_X0", "adamw")
        new = mutate_choice(rng, ["adamw", "saliencew"], current_resid)
        return Candidate(family, "scalar_optimizer", "Partition", f"set scalar optimizers to {new}", {
            **base_overrides,
            "AUTORESEARCH_ADAMW_VARIANT_RESID": new,
            "AUTORESEARCH_ADAMW_VARIANT_X0": new if current_x0 == current_resid else new,
        })
    if family == "value_embeds_betas":
        current = get_override(base_overrides, "AUTORESEARCH_VALUE_EMBEDS_BETAS", "0.8,0.95")
        new = mutate_choice(rng, ["0.8,0.95", "0.7,0.92", "0.75,0.93", "0.85,0.96"], current)
        return Candidate(family, "value_embedding_optimizer", "Gate", f"set value_embeds betas to {new}", {
            **base_overrides,
            "AUTORESEARCH_VALUE_EMBEDS_BETAS": new,
        })
    if family == "value_embeds_eps":
        current = get_override(base_overrides, "AUTORESEARCH_VALUE_EMBEDS_EPS", "1e-10")
        new = mutate_choice(rng, ["1e-10", "1e-9", "1e-8", "1e-7"], current)
        return Candidate(family, "value_embedding_optimizer", "Gate", f"set value_embeds eps to {new}", {
            **base_overrides,
            "AUTORESEARCH_VALUE_EMBEDS_EPS": new,
        })
    if family == "muon_gate":
        w_n = rng.choice(["0.14", "0.18", "0.22", "0.26"])
        w_a = rng.choice(["0.00", "0.02", "0.04"])
        w_f = rng.choice(["0.00", "0.02", "0.04"])
        gate_max = rng.choice(["1.08", "1.12", "1.16"])
        return Candidate(family, "muon_lr_gate", "Gate", f"muon gate N={w_n} A={w_a} F={w_f} max={gate_max}", {
            **base_overrides,
            "AUTORESEARCH_MUON_SALIENCE_W_N": w_n,
            "AUTORESEARCH_MUON_SALIENCE_W_A": w_a,
            "AUTORESEARCH_MUON_SALIENCE_W_F": w_f,
            "AUTORESEARCH_MUON_SALIENCE_GATE_MAX": gate_max,
        })
    if family == "saliencew_gains":
        p_gain = rng.choice(["0.12", "0.18", "0.24"])
        q_gain = rng.choice(["0.02", "0.04", "0.06"])
        xi_gain = rng.choice(["0.02", "0.04", "0.06"])
        gate_min = rng.choice(["0.90", "0.95", "0.98"])
        gate_max = rng.choice(["1.10", "1.15", "1.20"])
        return Candidate(family, "adamw_second_moment", "Gate", f"saliencew gains P={p_gain} Q={q_gain} Xi={xi_gain}", {
            **base_overrides,
            "AUTORESEARCH_SALIENCEW_PROMOTIVE_GAIN": p_gain,
            "AUTORESEARCH_SALIENCEW_AVERSIVE_GAIN": q_gain,
            "AUTORESEARCH_SALIENCEW_CONFLICT_GAIN": xi_gain,
            "AUTORESEARCH_SALIENCEW_GATE_MIN": gate_min,
            "AUTORESEARCH_SALIENCEW_GATE_MAX": gate_max,
        })
    if family == "schedule":
        matrix_lr = rng.choice(["0.04", "0.05", "0.06"])
        warmdown = rng.choice(["0.05", "0.10", "0.15"])
        final_frac = rng.choice(["0.0", "0.05", "0.1"])
        return Candidate(family, "schedule", "Gate", f"matrix_lr={matrix_lr} warmdown={warmdown} final_lr_frac={final_frac}", {
            **base_overrides,
            "AUTORESEARCH_MATRIX_LR": matrix_lr,
            "AUTORESEARCH_WARMDOWN_RATIO": warmdown,
            "AUTORESEARCH_FINAL_LR_FRAC": final_frac,
        })
    raise ValueError(f"Unknown family {family}")


def run_train(cwd: Path, regime, log_path: Path, overrides: dict[str, str] | None = None, timeout_seconds: float | None = None):
    env = os.environ.copy()
    if overrides:
        env.update(overrides)
    cmd = build_train_command(cwd, regime)
    with log_path.open("w", encoding="utf-8") as handle:
        proc = run_with_tree_timeout(cmd, cwd=cwd, timeout_seconds=timeout_seconds, stdout=handle, stderr=subprocess.STDOUT)
        if proc == 124:
            handle.write(f"\nTIMEOUT after {timeout_seconds}s\n")
    checkpoint_path = cwd / "checkpoint_pre_eval.pt"
    if checkpoint_path.exists():
        checkpoint_path.unlink()
    return proc


def write_best(path: Path, commit: str, regime: str, decision_mode: str, candidate: Candidate, metrics: dict, log_path: Path):
    payload = {
        "commit": commit,
        "regime": regime,
        "decision_mode": decision_mode,
        "description": candidate.note,
        "family": candidate.family,
        "part": candidate.part,
        "class": candidate.cls,
        "proposal_mode": candidate.proposal_mode,
        "log_path": str(log_path),
        "overrides": candidate.overrides,
        "eval_tokens": metrics.get("eval_tokens"),
        "memory_gb": metrics.get("memory_gb"),
        "num_steps": metrics.get("num_steps"),
        "peak_vram_mb": metrics.get("peak_vram_mb"),
        "total_seconds": metrics.get("total_seconds"),
        "train_batch_size": metrics.get("train_batch_size"),
        "training_seconds": metrics.get("training_seconds"),
        "val_bpb": metrics.get("val_bpb"),
    }
    path.write_text(json.dumps(payload, indent=2, sort_keys=True), encoding="utf-8")


def append_row(path: Path, header: str, values: list[str]):
    ensure_tsv(path, header)
    with path.open("a", encoding="utf-8") as handle:
        handle.write("\t".join(values) + "\n")


def git_head_short():
    return subprocess.check_output(["git", "rev-parse", "--short", "HEAD"], cwd=ROOT, text=True).strip()


def promote_candidate_code(candidate_dir: Path):
    champion_train = ROOT / "train.py"
    candidate_train = candidate_dir / "train.py"
    champion_train.write_text(candidate_train.read_text(encoding="utf-8"), encoding="utf-8")


def run_proxy_iteration(candidate: Candidate, run_id: str, tie_band: float, max_total_seconds: float, cwd: Path, timeout_seconds: float):
    log_path = PROXY_RUNS_DIR / f"{run_id}_{candidate.family}.log"
    rc = run_train(cwd, "proxy_fast", log_path, candidate.overrides, timeout_seconds=timeout_seconds)
    if rc != 0:
        return {"status": "crash", "log_path": log_path, "metrics": None, "decision": None}
    metrics = parse_run_log(log_path)
    reference = load_reference(PROXY_BEST_PATH)
    decision = decide(metrics, reference, mode="runtime", tie_band=tie_band, max_total_seconds=max_total_seconds)
    return {"status": "ok", "log_path": log_path, "metrics": metrics, "decision": decision}


def maybe_run_truth_promotion(candidate: Candidate, run_id: str, tie_band: float, max_total_seconds: float, cwd: Path, timeout_seconds: float):
    ref_path = TRUTH_OVERRIDE_BEST_PATH if TRUTH_OVERRIDE_BEST_PATH.exists() else TRUTH_REFERENCE_PATH
    log_path = TRUTH_RUNS_DIR / f"{run_id}_{candidate.family}.log"
    rc = run_train(cwd, "truth", log_path, candidate.overrides, timeout_seconds=timeout_seconds)
    if rc != 0:
        return {"status": "crash", "log_path": log_path, "metrics": None, "decision": None, "reference": str(ref_path)}
    metrics = parse_run_log(log_path)
    reference = load_reference(ref_path)
    decision = decide(metrics, reference, mode="trunk", tie_band=tie_band, max_total_seconds=max_total_seconds)
    return {"status": "ok", "log_path": log_path, "metrics": metrics, "decision": decision, "reference": str(ref_path)}


def main():
    parser = argparse.ArgumentParser(description="Run the AlphaEvolve-style overnight proxy search loop.")
    parser.add_argument("--iterations", type=int, default=0, help="Maximum proxy iterations. 0 means forever.")
    parser.add_argument("--sleep-seconds", type=float, default=0.0, help="Sleep between proxy iterations.")
    parser.add_argument("--truth-cadence", type=int, default=0, help="Run one truth promotion every N proxy keeps. 0 disables.")
    parser.add_argument(
        "--proposal-engine",
        choices=("codex_patch", "codex", "bandit", "hybrid", "hybrid_patch"),
        default="codex",
        help="How to generate the next candidate. `hybrid_patch` falls back from code-edit evolution to bandit overrides.",
    )
    parser.add_argument("--proxy-tie-band", type=float, default=0.005, help="Tie band for proxy runtime decisions.")
    parser.add_argument("--proxy-max-total-seconds", type=float, default=180.0, help="Hard wall-clock cutoff for proxy runs.")
    parser.add_argument("--truth-tie-band", type=float, default=0.002, help="Tie band for truth trunk decisions.")
    parser.add_argument("--truth-max-total-seconds", type=float, default=900.0, help="Hard wall-clock cutoff for truth promotions.")
    parser.add_argument("--codex-timeout-seconds", type=float, default=180.0, help="Hard timeout for one Codex proposal.")
    parser.add_argument("--proxy-train-timeout-seconds", type=float, default=300.0, help="Hard timeout for one proxy_fast training subprocess.")
    parser.add_argument("--truth-train-timeout-seconds", type=float, default=1200.0, help="Hard timeout for one truth training subprocess.")
    parser.add_argument("--seed", type=int, default=42, help="RNG seed.")
    args = parser.parse_args()

    rng = random.Random(args.seed)
    PROXY_RUNS_DIR.mkdir(exist_ok=True)
    TRUTH_RUNS_DIR.mkdir(exist_ok=True)
    CANDIDATE_WORKTREES_DIR.mkdir(exist_ok=True)
    cleanup_stale_candidate_worktrees()
    state = load_state()
    source_commit = git_head_short()

    proxy_header = "run_id\tstarted_at\tsource_commit\tfamily\tpart\tclass\tstatus\tdecision\tval_bpb\ttotal_seconds\tnum_steps\tmemory_gb\tnote\toverrides_json\tlog_path"
    truth_header = "run_id\tstarted_at\tsource_commit\tfamily\tstatus\tdecision\tval_bpb\ttotal_seconds\tnum_steps\tmemory_gb\treference\tlog_path\tnote\toverrides_json"

    iteration = 0
    while True:
        if args.iterations and iteration >= args.iterations:
            break

        best_proxy = load_reference(PROXY_BEST_PATH)
        base_overrides = dict(best_proxy.get("overrides", {}))
        run_id = f"{datetime.now().strftime('%Y%m%d_%H%M%S')}_{state['iteration']:05d}"
        started_at = now_utc()
        candidate = None
        proposal_engine_used = args.proposal_engine
        candidate_cwd = ROOT
        candidate_worktree = None
        if args.proposal_engine in {"codex_patch", "hybrid_patch"}:
            try:
                recent_rows = load_recent_proxy_rows()
                print("  codex_patch: starting proposal", flush=True)
                candidate, candidate_worktree, _, _, _ = propose_code_patch_codex(
                    best_proxy,
                    recent_rows,
                    run_id,
                    args.codex_timeout_seconds,
                )
                print("  codex_patch: proposal ready", flush=True)
                candidate_cwd = candidate_worktree
            except Exception:
                if args.proposal_engine == "codex_patch":
                    raise
                proposal_engine_used = "codex"
        if candidate is None and proposal_engine_used in {"codex", "hybrid"}:
            try:
                recent_rows = load_recent_proxy_rows()
                print("  codex: starting proposal", flush=True)
                candidate = propose_candidate_codex(best_proxy, recent_rows, run_id, args.codex_timeout_seconds)
                print("  codex: proposal ready", flush=True)
            except Exception:
                if args.proposal_engine == "codex":
                    raise
                proposal_engine_used = "bandit"
        if candidate is None:
            family = choose_family(state)
            candidate = propose_candidate(family, base_overrides, rng)
        else:
            family = candidate.family
        print(f"[{started_at}] proxy {run_id} family={candidate.family} note={candidate.note}", flush=True)
        print(f"  proposal_engine={proposal_engine_used}", flush=True)

        print("  proxy_fast: starting eval", flush=True)
        proxy_result = run_proxy_iteration(
            candidate,
            run_id,
            args.proxy_tie_band,
            args.proxy_max_total_seconds,
            candidate_cwd,
            args.proxy_train_timeout_seconds,
        )
        state["iteration"] += 1
        stats = state["family_stats"].setdefault(family, {"trials": 0, "keeps": 0})
        stats["trials"] += 1

        if proxy_result["status"] == "ok":
            metrics = proxy_result["metrics"]
            decision = proxy_result["decision"]["decision"]
            print(
                f"  proxy result decision={decision} val_bpb={metrics['val_bpb']:.6f} "
                f"total_seconds={metrics['total_seconds']:.1f} num_steps={metrics['num_steps']}",
                flush=True,
            )
            if decision == "keep":
                stats["keeps"] += 1
                state["proxy_keeps"] += 1
                if candidate.proposal_mode == "code_patch" and candidate_worktree is not None:
                    promote_candidate_code(candidate_worktree)
                write_best(PROXY_BEST_PATH, source_commit, "proxy_fast", "runtime", candidate, metrics, proxy_result["log_path"])
            append_row(
                PROXY_RESULTS_PATH,
                proxy_header,
                [
                    run_id,
                    started_at,
                    source_commit,
                    candidate.family,
                    candidate.part,
                    candidate.cls,
                    "ok",
                    decision,
                    f"{metrics['val_bpb']:.6f}",
                    f"{metrics['total_seconds']:.1f}",
                    str(metrics["num_steps"]),
                    f"{metrics['memory_gb']:.3f}",
                    candidate.note,
                    json.dumps(candidate.overrides, sort_keys=True),
                    str(proxy_result["log_path"]),
                ],
            )

            if args.truth_cadence > 0 and decision == "keep" and state["proxy_keeps"] % args.truth_cadence == 0:
                print(f"  promoting {run_id} to truth", flush=True)
                truth_result = maybe_run_truth_promotion(
                    candidate,
                    run_id,
                    args.truth_tie_band,
                    args.truth_max_total_seconds,
                    candidate_cwd,
                    args.truth_train_timeout_seconds,
                )
                if truth_result["status"] == "ok":
                    truth_metrics = truth_result["metrics"]
                    truth_decision = truth_result["decision"]["decision"]
                    print(
                        f"  truth result decision={truth_decision} val_bpb={truth_metrics['val_bpb']:.6f} "
                        f"total_seconds={truth_metrics['total_seconds']:.1f} num_steps={truth_metrics['num_steps']}",
                        flush=True,
                    )
                    if truth_decision == "keep":
                        write_best(
                            TRUTH_OVERRIDE_BEST_PATH,
                            source_commit,
                            "truth",
                            "trunk",
                            candidate,
                            truth_metrics,
                            truth_result["log_path"],
                        )
                    append_row(
                        TRUTH_RESULTS_PATH,
                        truth_header,
                        [
                            run_id,
                            started_at,
                            source_commit,
                            candidate.family,
                            "ok",
                            truth_decision,
                            f"{truth_metrics['val_bpb']:.6f}",
                            f"{truth_metrics['total_seconds']:.1f}",
                            str(truth_metrics["num_steps"]),
                            f"{truth_metrics['memory_gb']:.3f}",
                            truth_result["reference"],
                            str(truth_result["log_path"]),
                            candidate.note,
                            json.dumps(candidate.overrides, sort_keys=True),
                        ],
                    )
                else:
                    print("  truth promotion crashed", flush=True)
                    append_row(
                        TRUTH_RESULTS_PATH,
                        truth_header,
                        [
                            run_id,
                            started_at,
                            source_commit,
                            candidate.family,
                            "crash",
                            "",
                            "",
                            "",
                            "",
                            "",
                            truth_result["reference"],
                            str(truth_result["log_path"]),
                            candidate.note,
                            json.dumps(candidate.overrides, sort_keys=True),
                        ],
                    )
        else:
            print(f"  proxy run crashed: {proxy_result['log_path']}", flush=True)
            append_row(
                PROXY_RESULTS_PATH,
                proxy_header,
                [
                    run_id,
                    started_at,
                    source_commit,
                    candidate.family,
                    candidate.part,
                    candidate.cls,
                    "crash",
                    "",
                    "",
                    "",
                    "",
                    "",
                    candidate.note,
                    json.dumps(candidate.overrides, sort_keys=True),
                    str(proxy_result["log_path"]),
                ],
            )

        if candidate_worktree is not None:
            remove_candidate_worktree(candidate_worktree)

        save_state(state)
        iteration += 1
        if args.sleep_seconds > 0:
            time.sleep(args.sleep_seconds)


if __name__ == "__main__":
    main()
