#!/usr/bin/env python3
"""
v2 hill-climb orchestrator for the pelican-bicycle-svg benchmark, targeting the
`jettygrowthteam` collection (all models routed through OpenRouter).

Adapted from jetty/pelican-bicycle-svg/orchestrator/hill_climb.py:
  * collection jettygrowthteam, task pelican-bicycle-svg (created by --create-task)
  * 5 rounds per agent (default), agents run in parallel threads
  * runs are launched multipart (init_params=<json>) exactly as a reader would with curl
  * finals come from the trajectory's /download zip
  * a shared spend ledger (ledger.json) stops new launches once estimated spend > --budget

Each round = one Jetty trajectory running the runbook with 3 inner self-critique rounds.
Between rounds the runbook is rewritten to embed the prior final.svg and target the
weakest self-scored axis (same surgery as v1).

Usage:
  python3 hill_climb_v2.py --create-task
  python3 hill_climb_v2.py --agent all --rounds 5
  python3 hill_climb_v2.py --agent opus --rounds 5
"""
import argparse, io, json, os, re, ssl, sys, threading, time, urllib.error, urllib.request, uuid, zipfile
from pathlib import Path

try:
    import certifi
    _SSL = ssl.create_default_context(cafile=certifi.where())
except ImportError:
    _SSL = ssl.create_default_context()

API = "https://flows-api.jetty.io"
COLLECTION = os.environ.get("JETTY_COLLECTION", "jettygrowthteam")
TASK = "pelican-bicycle-svg"
ROOT = Path(__file__).resolve().parent
BASELINE = ROOT / "runbook.md"
DL = ROOT / "downloads"
LEDGER = ROOT / "ledger.json"
ENV_FILE = Path(os.environ.get("JETTY_ENV_FILE", Path.home() / "Projects/jetty/.env"))  # optional fallback

# OpenRouter $/token (prompt, completion, cache_read, cache_write) checked 2026-09-27.
# jev-router is a router with variable pricing (-1): we charge it at Opus rates to stay conservative.
PRICES = {
    "anthropic/claude-sonnet-5":   (2e-6, 10e-6, 0.2e-6, 2.5e-6),
    "anthropic/claude-opus-5.5":   (4e-6, 20e-6, 0.2e-6, 5e-6),
    "google/gemini-3.8-flash":     (0.75e-6, 3.75e-6, 0.075e-6, 0.75e-6),
    "openai/gpt-5.6-terra":        (2e-6, 12e-6, 0.2e-6, 2.5e-6),
    "typesafe/jev-router":         (4e-6, 20e-6, 0.4e-6, 5e-6),
}

AGENTS = {
    "sonnet5":   {"agent": "claude-code", "model": "anthropic/claude-sonnet-5", "slug": "anthropic/claude-sonnet-5",
                  "label": "Claude Code · Sonnet 5", "dir": "claude-sonnet5"},
    "opus55":    {"agent": "claude-code", "model": "anthropic/claude-opus-5.5", "slug": "anthropic/claude-opus-5.5",
                  "label": "Claude Code · Opus 5.5", "dir": "claude-opus55"},
    "gemini38":  {"agent": "opencode", "model": "openrouter/google/gemini-3.8-flash", "slug": "google/gemini-3.8-flash",
                  "label": "OpenCode · Gemini 3.8 Flash", "dir": "opencode-gemini38"},
    "gpt56":     {"agent": "opencode", "model": "openrouter/openai/gpt-5.6-terra", "slug": "openai/gpt-5.6-terra",
                  "label": "OpenCode · GPT-5.6 Terra", "dir": "opencode-gpt56"},
    "jev":       {"agent": "hermes", "model": "openrouter/typesafe/jev-router", "slug": "typesafe/jev-router",
                  "label": "Hermes · Jev Router", "dir": "hermes-jev"},
}

TASK_WORKFLOW = {
    "init_params": {
        "agent": "claude-code",
        "model": "anthropic/claude-sonnet-5",
        "model_provider": "openrouter",
        "snapshot": "python312-uv",
        "instruction": None,  # filled with runbook.md
        "vars": {"prompt": "Execute the runbook end-to-end.", "results_dir": "/app/results", "max_rounds": 3},
        "file_paths": [],
    },
    "steps": ["run"],
    "step_configs": {"run": {
        "activity": "runbook",
        "agent_path": "init_params.agent",
        "model_path": "init_params.model",
        "model_provider_path": "init_params.model_provider",
        "snapshot_path": "init_params.snapshot",
        "instruction_path": "init_params.instruction",
        "template_variables_path": "init_params.vars",
        "files_path": "init_params.file_paths",
        "cpus": 4, "memory": "8G", "timeout_sec": 1800, "network_enabled": True,
    }},
}

HEADLESS_PROMPT = (
    "Execute the runbook end-to-end NOW using your file and shell tools; actually write the files, "
    "do not just describe the steps. max_rounds={max_rounds}, results_dir=/app/results. "
    "You are running FULLY HEADLESS: there is no human to answer questions, so NEVER stop to ask for "
    "clarification or to choose between options; pick the best path and keep going. If you cannot view "
    "the rendered PNG as an image, do NOT halt: self-critique by reasoning over the SVG source and its "
    "coordinate geometry instead, assign the four axis scores on that basis, and continue through every round. "
    "You MUST finish by writing ALL THREE of these files to /app/results/: final.svg (the best round's SVG), "
    "final.png (its rasterized render), and report.md. report.md MUST contain a line with the text "
    "'Per-round scores' followed by a markdown table whose header is exactly "
    "'| Round | Pelican | Bicycle | Composition | Polish | Total |' and one row per round (each axis out of 10, "
    "Total out of 40). Do not end your turn until all three files exist on disk."
)


def token():
    t = os.environ.get("JETTY_TOKEN")
    if t:
        return t.strip()
    if ENV_FILE.exists():
        for line in ENV_FILE.read_text().splitlines():
            if line.startswith("JETTY_API_TOKEN_JETTY_GROWTHTEAM="):
                return line.split("=", 1)[1].strip().strip('"').strip("'")
    sys.exit("Set JETTY_TOKEN to a Jetty API key for your collection (and JETTY_COLLECTION to its name).")


def _req(method, path, body=None, headers=None, timeout=60, raw=False):
    h = {"Authorization": f"Bearer {token()}", "Accept": "application/json", "User-Agent": "evaljetty-pelicans-v2/1.0"}
    h.update(headers or {})
    req = urllib.request.Request(API + path, data=body, headers=h, method=method)
    try:
        with urllib.request.urlopen(req, timeout=timeout, context=_SSL) as r:
            data = r.read()
            return data if raw else json.loads(data)
    except urllib.error.HTTPError as e:
        raise RuntimeError(f"{method} {path} -> HTTP {e.code}: {e.read().decode('utf-8', 'replace')[:600]}") from None


def api_json(method, path, obj, timeout=60):
    return _req(method, path, json.dumps(obj).encode(), {"Content-Type": "application/json"}, timeout)


def api_multipart(path, fields, timeout=120):
    b = uuid.uuid4().hex
    parts = []
    for k, v in fields.items():
        parts.append(f"--{b}\r\nContent-Disposition: form-data; name=\"{k}\"\r\n\r\n".encode() + v.encode() + b"\r\n")
    body = b"".join(parts) + f"--{b}--\r\n".encode()
    return _req("POST", path, body, {"Content-Type": f"multipart/form-data; boundary={b}"}, timeout)


# ---- runbook surgery (same as v1 orchestrator) ------------------------------------------
def replace_frontmatter(md, meta, round_n):
    m = re.match(r"^---\n(.*?)\n---\n", md, re.DOTALL)
    fm = m.group(1)
    for field, val in (("version", f"2.{round_n}"), ("agent", meta["agent"]), ("model", meta["model"]),
                       ("model_provider", "openrouter")):
        fm = re.sub(rf"^{field}:.*$", f"{field}: {val}", fm, count=1, flags=re.MULTILINE)
    return f"---\n{fm}\n---\n" + md[m.end():]


def embed_baseline_svg(md, svg):
    block = f"```svg\n<!-- BASELINE SVG: paste byte-for-byte into rounds/v1.svg -->\n{svg.strip()}\n```"
    m = re.search(r"```svg\s*\n(.*?)\n```", md, re.DOTALL)
    if m:
        return md[:m.start()] + block + md[m.end():]
    marker = "### 2. Round 1"
    idx = md.index(marker)
    end = md.find("\n\n", idx)
    lead = ("\n\n**This is a hill-climb round.** Start from the baseline below (the previous round's best SVG): "
            "write it byte-for-byte to `rounds/v1.svg`, score it, then improve on it. Make targeted edits, keep what works.")
    return md[:end] + lead + "\n\n" + block + md[end:]


WEAKEST = {
    "pelican": "Sharpen pelican anatomy: pouch droop more visible, eye placement higher, wing-tip definition crisper. Target Pelican >= 9.",
    "bicycle": "Tighten bicycle geometry: both wheels perfect circles, visible spokes, frame uses a non-orange color so legs aren't confused with frame. Target Bicycle >= 9.",
    "composition": "Improve composition: close any wing-to-grip gap, ensure body sits over saddle (saddle visible), feet land ON pedals. Target Composition >= 10.",
    "polish": "Polish: line quality, no overlapping clutter, consistent gradient direction, subtle ground shadow. Target Polish >= 9.",
}


def update_description(md, axis, n):
    d = f"Round {n}: iterate on prior final.svg. {WEAKEST[axis]} Regression guard: revert if a round scores lower."
    return re.sub(r"^description:.*$", f"description: {d}", md, count=1, flags=re.MULTILINE)


AXES = ["pelican", "bicycle", "composition", "polish"]


def parse_round_scores(rep):
    out, on = [], False
    for line in rep.splitlines():
        if "Per-round scores" in line:
            on = True
            continue
        if on and line.startswith("|") and not re.match(r"^\|\s*-", line) and "Round" not in line:
            c = [x.strip().strip("*") for x in line.strip().strip("|").split("|")]
            try:
                row = [float(re.findall(r"[\d.]+", x)[0]) for x in c[1:6]]
                out.append(dict(zip(AXES + ["total"], row)))
            except Exception:
                pass
        elif on and out and not line.strip():
            break
    return out


# ---- ledger ---------------------------------------------------------------------------
_lock = threading.Lock()


def ledger():
    return json.loads(LEDGER.read_text()) if LEDGER.exists() else {"runs": {}}


def ledger_put(tid, rec):
    with _lock:
        L = ledger()
        L["runs"][tid] = rec
        L["total_usd"] = round(sum(r.get("cost_usd") or r.get("est_usd", 0) for r in L["runs"].values()), 4)
        LEDGER.write_text(json.dumps(L, indent=2))
        return L["total_usd"]


def spend():
    L = ledger()
    return sum((r.get("cost_usd") or r.get("est_usd", 0)) for r in L["runs"].values())


# Fallback when the runner reports no token usage (opencode/hermes sometimes don't): conservative flat estimate.
FALLBACK_USD = {"anthropic/claude-sonnet-5": 1.5, "anthropic/claude-opus-5.5": 2.5, "google/gemini-3.8-flash": 0.5,
                "openai/gpt-5.6-terra": 1.5, "typesafe/jev-router": 2.5}


def estimate_cost(slug, usage):
    if not usage or not (usage.get("total_tokens") or usage.get("completion_tokens")):
        return FALLBACK_USD[slug]
    if usage.get("cost_usd"):
        return float(usage["cost_usd"])
    p, c, cr, cw = PRICES[slug]
    return (usage.get("prompt_tokens", 0) or 0) * p + (usage.get("completion_tokens", 0) or 0) * c + \
        (usage.get("cache_read_input_tokens", 0) or 0) * cr + (usage.get("cache_creation_input_tokens", 0) or 0) * cw


# ---- run plumbing ---------------------------------------------------------------------
def submit(md, meta, max_rounds=3):
    ip = {"agent": meta["agent"], "model": meta["model"], "model_provider": "openrouter", "instruction": md,
          "vars": {"prompt": HEADLESS_PROMPT.format(max_rounds=max_rounds), "results_dir": "/app/results",
                   "max_rounds": max_rounds}}
    r = api_multipart(f"/api/v1/run/{COLLECTION}/{TASK}", {"init_params": json.dumps(ip)})
    tid = r.get("trajectory_id") or r.get("workflow_id", "").split("--")[-1][:8]
    if not tid:
        raise RuntimeError(f"no trajectory id in {r}")
    return tid


def get_traj(tid):
    return _req("GET", f"/api/v1/db/trajectory/{COLLECTION}/{TASK}/{tid}", timeout=30)


def poll(tid, tag, timeout=2700, every=30):
    t0 = time.time()
    while time.time() - t0 < timeout:
        try:
            d = get_traj(tid)
        except Exception as e:
            print(f"[{tag}] poll err {e}", flush=True)
            time.sleep(every)
            continue
        st = (d.get("status") or "").lower()
        if st in ("completed", "success", "done", "failed", "error", "cancelled"):
            return d
        time.sleep(every)
    raise RuntimeError(f"{tid} timed out")


def download(tid, out):
    data = _req("GET", f"/api/v1/trajectory/{COLLECTION}/{TASK}/{tid}/download", timeout=300, raw=True)
    (out / "trajectory.zip").write_bytes(data)
    saved = []
    with zipfile.ZipFile(io.BytesIO(data)) as z:
        for n in z.namelist():
            m = re.search(r"app/results/(.+)$", n) or re.search(r"app--results--(.+)$", n)
            if not m:
                a = re.search(r"/agent/(.+\.(?:txt|log))$", n)
                if a:
                    t = out / "agent" / a.group(1).replace("/", "_")
                    t.parent.mkdir(parents=True, exist_ok=True)
                    t.write_bytes(z.read(n))
                continue
            local = m.group(1).replace("--", "/")
            if local in ("final.svg", "final.png", "report.md") or local.startswith("rounds/"):
                t = out / local
                t.parent.mkdir(parents=True, exist_ok=True)
                t.write_bytes(z.read(n))
                saved.append(local)
    return saved


def run_agent(key, rounds, budget, max_rounds=3):
    meta = AGENTS[key]
    tag = key
    out_root = DL / meta["dir"]
    out_root.mkdir(parents=True, exist_ok=True)
    md = update_description(replace_frontmatter(BASELINE.read_text(), meta, 1), "composition", 1)
    prior_svg = prior = None
    failures = 0
    n = 1
    while n <= rounds:
        rd = out_root / f"v{n}"
        rd.mkdir(exist_ok=True)
        if (rd / "final.svg").exists() and (rd / "report.md").exists():
            sc = parse_round_scores((rd / "report.md").read_text())
            prior = max(sc, key=lambda r: r["total"]) if sc else prior
            prior_svg = (rd / "final.svg").read_text()
            print(f"[{tag}] v{n} SKIP (exists)", flush=True)
            n += 1
            continue
        if prior_svg and prior:
            wa = min(AXES, key=lambda a: prior[a])
            md = replace_frontmatter(update_description(embed_baseline_svg(md, prior_svg), wa, n), meta, n)
        s = spend()
        if s > budget:
            print(f"[{tag}] BUDGET STOP at v{n}: est spend ${s:.2f} > ${budget}", flush=True)
            (out_root / "STOPPED").write_text(f"budget stop before v{n}: ${s:.2f}\n")
            return
        tid = None
        if (rd / "trajectory_id").exists():  # reattach to an in-flight/finished run after an orchestrator restart
            tid = (rd / "trajectory_id").read_text().strip()
            if (rd / f"failed_{tid}").exists():
                tid = None
            else:
                print(f"[{tag}] v{n} reattach {tid}", flush=True)
        if not tid:
            (rd / "runbook.md").write_text(md)
        try:
            tid = tid or submit(md, meta, max_rounds)
        except Exception as e:
            failures += 1
            print(f"[{tag}] v{n} submit failed ({failures}): {e}", flush=True)
            if failures >= 2:
                (out_root / "DROPPED").write_text(f"submit failures: {e}\n")
                return
            time.sleep(20)
            continue
        (rd / "trajectory_id").write_text(tid)
        print(f"[{tag}] v{n} submitted {tid} https://jetty.io/{COLLECTION}/{TASK}/{tid}", flush=True)
        if tid not in ledger()["runs"] or ledger()["runs"][tid].get("status") == "running":
            ledger_put(tid, {"agent": key, "round": n, "status": "running", "est_usd": FALLBACK_USD[meta["slug"]]})
        d = poll(tid, tag)
        (rd / "trajectory.json").write_text(json.dumps({k: v for k, v in d.items()}, indent=1, default=str))
        run = (d.get("steps") or {}).get("run") or {}
        usage = (run.get("outputs") or {}).get("usage") or {}
        cost = estimate_cost(meta["slug"], usage)
        st = (d.get("status") or "").lower()
        total = ledger_put(tid, {"agent": key, "round": n, "status": st, "usage": usage,
                                 "est_usd": round(cost, 4) if cost is not None else 1.5,
                                 "duration_s": run.get("duration_seconds")})
        saved = []
        try:
            saved = download(tid, rd)
        except Exception as e:
            print(f"[{tag}] v{n} download err {e}", flush=True)
        print(f"[{tag}] v{n} {tid} status={st} cost~${cost if cost is None else round(cost, 2)} ledger=${total:.2f} saved={saved}", flush=True)
        if st != "completed" or not (rd / "final.svg").exists():
            failures += 1
            (rd / f"failed_{tid}").write_text(json.dumps({"status": st, "error": d.get("error")}, default=str))
            for f in ("trajectory_id",):
                (rd / f).unlink(missing_ok=True)
            print(f"[{tag}] v{n} FAILED ({failures}) err={str(d.get('error'))[:300]}", flush=True)
            if failures >= 2:
                (out_root / "DROPPED").write_text(f"failed twice; last {tid}: {d.get('error')}\n")
                return
            continue
        (rd / "trajectory_id").write_text(tid)
        if (rd / "report.md").exists():
            sc = parse_round_scores((rd / "report.md").read_text())
            if sc:
                prior = max(sc, key=lambda r: r["total"])
        prior_svg = (rd / "final.svg").read_text()
        n += 1


def create_task():
    wf = json.loads(json.dumps(TASK_WORKFLOW))
    wf["init_params"]["instruction"] = BASELINE.read_text()
    (ROOT / "task_workflow.json").write_text(json.dumps(wf, indent=2))
    body = {"name": TASK, "description": "Pelican riding a bicycle: hand-written SVG with 3 self-critique rounds "
                                         "(evaljetty.com/pelicans v2 refresh).", "workflow": wf}
    try:
        r = api_json("POST", f"/api/v1/tasks/{COLLECTION}", body)
    except RuntimeError as e:
        if "exist" in str(e).lower() or "409" in str(e):
            r = api_json("PUT", f"/api/v1/tasks/{COLLECTION}/{TASK}", {"workflow": wf, "description": body["description"]})
        else:
            raise
    print(json.dumps(r)[:400])


def main():
    ap = argparse.ArgumentParser()
    ap.add_argument("--agent", default="all")
    ap.add_argument("--rounds", type=int, default=5)
    ap.add_argument("--budget", type=float, default=60.0)
    ap.add_argument("--create-task", action="store_true")
    a = ap.parse_args()
    if a.create_task:
        create_task()
        return
    keys = list(AGENTS) if a.agent == "all" else a.agent.split(",")
    ths = [threading.Thread(target=run_agent, args=(k, a.rounds, a.budget), name=k) for k in keys]
    for t in ths:
        t.start()
        time.sleep(3)
    for t in ths:
        t.join()
    print(f"ALL DONE. est spend ${spend():.2f}")


if __name__ == "__main__":
    main()
