import argparse
import json
import math
import re
from pathlib import Path


SUMMARY_PATTERNS = {
    "regime": re.compile(r"regime:\s+([A-Za-z0-9_-]+)"),
    "val_bpb": re.compile(r"val_bpb:\s+([0-9.]+)"),
    "training_seconds": re.compile(r"training_seconds:\s+([0-9.]+)"),
    "total_seconds": re.compile(r"total_seconds:\s+([0-9.]+)"),
    "peak_vram_mb": re.compile(r"peak_vram_mb:\s+([0-9.]+)"),
    "total_tokens_M": re.compile(r"total_tokens_M:\s+([0-9.]+)"),
    "eval_tokens": re.compile(r"eval_tokens:\s+([0-9]+)"),
    "num_steps": re.compile(r"num_steps:\s+([0-9]+)"),
    "num_params_M": re.compile(r"num_params_M:\s+([0-9.]+)"),
    "depth": re.compile(r"depth:\s+([0-9]+)"),
    "aspect_ratio": re.compile(r"aspect_ratio:\s+([0-9]+)"),
    "total_batch_size": re.compile(r"total_batch_size:\s+([0-9]+)"),
    "warmup_excluded_steps": re.compile(r"warmup_excluded_steps:\s+([0-9]+)"),
    "train_batch_size": re.compile(r"train_batch_size:\s+([0-9]+)"),
    "eval_batch_size": re.compile(r"eval_batch_size:\s+([0-9]+)"),
    "activation_checkpointing": re.compile(r"activation_checkpointing:\s+(enabled|disabled)"),
}


def parse_run_log(path: Path):
    raw = path.read_bytes()
    if raw.startswith(b"\xff\xfe") or raw.startswith(b"\xfe\xff"):
        text = raw.decode("utf-16")
    else:
        text = raw.decode("utf-8", errors="replace")
    metrics = {}
    for line in text.splitlines():
        stripped = line.strip()
        for key, pattern in SUMMARY_PATTERNS.items():
            match = pattern.search(stripped)
            if match:
                value = match.group(1)
                if key in {
                    "eval_tokens",
                    "num_steps",
                    "depth",
                    "aspect_ratio",
                    "total_batch_size",
                    "warmup_excluded_steps",
                    "train_batch_size",
                    "eval_batch_size",
                }:
                    metrics[key] = int(value)
                elif key in {"activation_checkpointing", "regime"}:
                    metrics[key] = value
                else:
                    metrics[key] = float(value)
    required = {"val_bpb", "total_seconds", "peak_vram_mb", "num_steps"}
    missing = sorted(required - metrics.keys())
    if missing:
        raise ValueError(f"Missing required metrics in {path}: {', '.join(missing)}")
    metrics["memory_gb"] = round(metrics["peak_vram_mb"] / 1024.0, 3)
    if "total_tokens_M" in metrics and metrics["total_seconds"] > 0:
        metrics["effective_tok_per_sec"] = round(metrics["total_tokens_M"] * 1_000_000 / metrics["total_seconds"], 1)
    else:
        metrics["effective_tok_per_sec"] = None
    return metrics


def load_reference(path: Path):
    return json.loads(path.read_text(encoding="utf-8"))


def decide(run_metrics, ref_metrics, mode, tie_band, max_total_seconds):
    reasons = []
    if run_metrics.get("regime") and ref_metrics.get("regime") and run_metrics["regime"] != ref_metrics["regime"]:
        reasons.append(f"regime mismatch: run={run_metrics['regime']} ref={ref_metrics['regime']}")
        return {"decision": "discard", "reasons": reasons}
    if run_metrics["total_seconds"] > max_total_seconds:
        reasons.append(f"hard fail: total_seconds {run_metrics['total_seconds']:.1f} > {max_total_seconds:.1f}")
        return {"decision": "discard", "reasons": reasons}

    delta_bpb = run_metrics["val_bpb"] - ref_metrics["val_bpb"]
    delta_seconds = run_metrics["total_seconds"] - ref_metrics["total_seconds"]
    delta_steps = run_metrics["num_steps"] - ref_metrics["num_steps"]
    delta_vram = run_metrics["peak_vram_mb"] - ref_metrics["peak_vram_mb"]

    reasons.append(f"delta_bpb={delta_bpb:+.6f}")
    reasons.append(f"delta_seconds={delta_seconds:+.1f}")
    reasons.append(f"delta_steps={delta_steps:+d}")
    reasons.append(f"delta_vram_mb={delta_vram:+.1f}")

    if delta_bpb <= -tie_band:
        reasons.append("material val_bpb improvement")
        return {"decision": "keep", "reasons": reasons}

    if abs(delta_bpb) <= tie_band:
        proxy_wins = 0
        if delta_seconds < 0:
            proxy_wins += 1
        if delta_steps > 0:
            proxy_wins += 1
        if delta_vram < 0:
            proxy_wins += 1
        if delta_seconds == 0 and delta_steps == 0 and delta_vram == 0:
            reasons.append("exact parity with reference")
            return {"decision": "keep", "reasons": reasons}
        reasons.append(f"near-tie on val_bpb, proxy_wins={proxy_wins}")
        return {"decision": "keep" if proxy_wins >= 2 else "discard", "reasons": reasons}

    if mode == "runtime" and delta_bpb <= 0.010:
        proxy_wins = 0
        if delta_seconds < -5.0:
            proxy_wins += 1
        if delta_steps > 0:
            proxy_wins += 1
        if delta_vram < -256.0:
            proxy_wins += 1
        reasons.append(f"runtime mode with bounded bpb regression, proxy_wins={proxy_wins}")
        return {"decision": "keep" if proxy_wins >= 2 else "discard", "reasons": reasons}

    reasons.append("val_bpb regression outside keep band")
    return {"decision": "discard", "reasons": reasons}


def main():
    parser = argparse.ArgumentParser(description="Score an autoresearch run under the layered metric policy.")
    parser.add_argument("run_log", help="Path to run.log")
    parser.add_argument("--reference", default="current_best_run.json", help="Reference run JSON path")
    parser.add_argument("--mode", choices=("trunk", "runtime"), default="trunk", help="Decision mode")
    parser.add_argument("--tie-band", type=float, default=0.002, help="Absolute val_bpb tie band")
    parser.add_argument("--max-total-seconds", type=float, default=600.0, help="Hard wall-clock cutoff")
    args = parser.parse_args()

    run_metrics = parse_run_log(Path(args.run_log))
    ref_metrics = load_reference(Path(args.reference))
    decision = decide(run_metrics, ref_metrics, args.mode, args.tie_band, args.max_total_seconds)

    payload = {
        "mode": args.mode,
        "reference": str(Path(args.reference)),
        "run": run_metrics,
        "reference_metrics": ref_metrics,
        "decision": decision["decision"],
        "reasons": decision["reasons"],
    }
    print(json.dumps(payload, indent=2, sort_keys=True))


if __name__ == "__main__":
    main()
