import argparse
import json
import math
import re
from pathlib import Path


SUMMARY_PATTERNS = {
    "run_mode": re.compile(r"run_mode:\s+([A-Za-z0-9_]+)"),
    "objective": re.compile(r"objective:\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.]+)"),
    "num_steps": re.compile(r"num_steps:\s+([0-9]+)"),
    "warmup_steps": re.compile(r"warmup_steps:\s+([0-9]+)"),
    "measured_steps": re.compile(r"measured_steps:\s+([0-9]+)"),
    "num_params_M": re.compile(r"num_params_M:\s+([0-9.]+)"),
    "depth": re.compile(r"depth:\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]+)"),
    "train_step_ms": re.compile(r"train_step_ms:\s+([0-9.]+)"),
    "forward_step_ms": re.compile(r"forward_step_ms:\s+([0-9.]+)"),
    "backward_step_ms": re.compile(r"backward_step_ms:\s+([0-9.]+)"),
    "optimizer_step_ms": re.compile(r"optimizer_step_ms:\s+([0-9.]+)"),
    "backprop_step_ms": re.compile(r"backprop_step_ms:\s+([0-9.]+)"),
    "backprop_share_percent": re.compile(r"backprop_share_percent:\s+([0-9.]+)"),
    "measured_tok_per_sec": re.compile(r"measured_tok_per_sec:\s+([0-9.]+)"),
    "activation_checkpointing": re.compile(r"activation_checkpointing:\s+(enabled|disabled)"),
}

INT_FIELDS = {"num_steps", "warmup_steps", "measured_steps", "depth", "train_batch_size", "eval_batch_size"}
STRING_FIELDS = {"run_mode", "objective", "activation_checkpointing"}
MODE_REQUIRED_FIELDS = {
    "trunk": {"val_bpb", "total_seconds", "peak_vram_mb", "num_steps"},
    "runtime": {"val_bpb", "total_seconds", "peak_vram_mb", "num_steps"},
    "backprop": {
        "backprop_step_ms",
        "train_step_ms",
        "forward_step_ms",
        "total_seconds",
        "peak_vram_mb",
        "num_steps",
        "measured_steps",
    },
}


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 INT_FIELDS:
                    metrics[key] = int(value)
                elif key in STRING_FIELDS:
                    metrics[key] = value
                else:
                    metrics[key] = float(value)
    return _normalize_metrics(metrics)


def _normalize_metrics(metrics):
    if "peak_vram_mb" in metrics and "memory_gb" not in metrics:
        metrics["memory_gb"] = round(metrics["peak_vram_mb"] / 1024.0, 3)
    if "total_tokens_M" in metrics and "total_seconds" in metrics and metrics["total_seconds"] > 0:
        metrics["effective_tok_per_sec"] = round(metrics["total_tokens_M"] * 1_000_000 / metrics["total_seconds"], 1)
    elif "effective_tok_per_sec" not in metrics:
        metrics["effective_tok_per_sec"] = None
    return metrics


def load_reference(path: Path):
    return _normalize_metrics(json.loads(path.read_text(encoding="utf-8")))


def require_metrics(metrics, mode, source_label):
    required = MODE_REQUIRED_FIELDS[mode]
    missing = sorted(required - metrics.keys())
    if missing:
        raise ValueError(f"Missing required {mode} metrics in {source_label}: {', '.join(missing)}")


def decide(run_metrics, ref_metrics, mode, tie_band, max_total_seconds):
    if mode == "backprop":
        return decide_backprop(run_metrics, ref_metrics, tie_band, max_total_seconds)
    return decide_trunk(run_metrics, ref_metrics, mode, tie_band, max_total_seconds)


def decide_trunk(run_metrics, ref_metrics, mode, tie_band, max_total_seconds):
    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 decide_backprop(run_metrics, ref_metrics, tie_band, max_total_seconds):
    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_backprop_ms = run_metrics["backprop_step_ms"] - ref_metrics["backprop_step_ms"]
    delta_train_step_ms = run_metrics["train_step_ms"] - ref_metrics["train_step_ms"]
    delta_forward_ms = run_metrics["forward_step_ms"] - ref_metrics["forward_step_ms"]
    delta_vram = run_metrics["peak_vram_mb"] - ref_metrics["peak_vram_mb"]

    reasons.append(f"delta_backprop_step_ms={delta_backprop_ms:+.3f}")
    reasons.append(f"delta_train_step_ms={delta_train_step_ms:+.3f}")
    reasons.append(f"delta_forward_step_ms={delta_forward_ms:+.3f}")
    reasons.append(f"delta_vram_mb={delta_vram:+.1f}")

    if delta_backprop_ms <= -tie_band:
        reasons.append("material backprop latency improvement")
        return {"decision": "keep", "reasons": reasons}

    if abs(delta_backprop_ms) <= tie_band:
        proxy_wins = 0
        if delta_train_step_ms < 0:
            proxy_wins += 1
        if delta_forward_ms < 0:
            proxy_wins += 1
        if delta_vram < 0:
            proxy_wins += 1
        if delta_train_step_ms == 0 and delta_forward_ms == 0 and delta_vram == 0:
            reasons.append("exact parity with reference")
            return {"decision": "keep", "reasons": reasons}
        reasons.append(f"near-tie on backprop_step_ms, proxy_wins={proxy_wins}")
        return {"decision": "keep" if proxy_wins >= 2 else "discard", "reasons": reasons}

    reasons.append("backprop latency 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", "backprop"), default="trunk", help="Decision mode")
    parser.add_argument(
        "--tie-band",
        type=float,
        default=None,
        help="Absolute tie band. Uses BPB units for trunk/runtime and milliseconds for backprop.",
    )
    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))
    require_metrics(run_metrics, args.mode, args.run_log)
    require_metrics(ref_metrics, args.mode, args.reference)

    if args.tie_band is None:
        tie_band = 1.0 if args.mode == "backprop" else 0.002
    else:
        tie_band = args.tie_band

    decision = decide(run_metrics, ref_metrics, args.mode, tie_band, args.max_total_seconds)

    payload = {
        "mode": args.mode,
        "reference": str(Path(args.reference)),
        "run": run_metrics,
        "reference_metrics": ref_metrics,
        "tie_band": tie_band,
        "decision": decision["decision"],
        "reasons": decision["reasons"],
    }
    print(json.dumps(payload, indent=2, sort_keys=True))


if __name__ == "__main__":
    main()
