#!/usr/bin/env python3
"""Export a grok session trace; if --from HOST, rsync the tarball into local /tmp.

HOST can be user@host (e.g. janitor@ravenrock). Use --jump when you need a
ProxyJump hop (e.g. Tailscale only reaches the bastion).
"""

import json
import subprocess
import sys
from argparse import ArgumentParser
from pathlib import Path
from tempfile import gettempdir

parser = ArgumentParser(description=__doc__)
parser.add_argument("session", help="grok session id")
parser.add_argument(
    "--from",
    dest="from_host",
    metavar="HOST",
    help="run on HOST via ssh (user@host ok) and rsync the tarball into local /tmp",
)
parser.add_argument(
    "-J",
    "--jump",
    metavar="JUMP",
    help="ssh ProxyJump hop (user@host ok); used for both ssh and rsync",
)
args = parser.parse_args()


def ssh_host(spec: str | None) -> str | None:
    """Host part of user@host (or bare host); None if unset."""
    if not spec:
        return None
    return spec.rsplit("@", 1)[-1]


if (h := ssh_host(args.jump)) is not None and h == ssh_host(args.from_host):
    args.jump = None

if args.jump and not args.from_host:
    parser.error("--jump only makes sense with --from")

TRACE_CMD = [
    "workspaced",
    "tool",
    "with",
    "grok-build",
    "--",
    "grok",
    "trace",
    "--local",
    "--json",
    args.session,
]


def die(msg: str, code: int = 1) -> None:
    print(f"error: {msg}", file=sys.stderr)
    raise SystemExit(code)


def ssh_base(from_host: str) -> list[str]:
    cmd = ["ssh"]
    if args.jump:
        cmd += ["-J", args.jump]
    cmd += [from_host, "--"]
    return cmd


def rsync_ssh() -> list[str]:
    """rsync -e 'ssh …' args so jumps match the ssh invocation."""
    if not args.jump:
        return []
    # single string for -e: rsync passes it to the shell-less exec of ssh
    return ["-e", f"ssh -J {args.jump}"]


def run_trace(from_host: str | None) -> Path:
    if from_host:
        via = f" via {args.jump}" if args.jump else ""
        print(f"exporting session {args.session} on {from_host}{via}…", file=sys.stderr)
        cmd = [*ssh_base(from_host), *TRACE_CMD]
    else:
        print(f"exporting session {args.session} locally…", file=sys.stderr)
        cmd = TRACE_CMD

    try:
        proc = subprocess.run(cmd, capture_output=True, check=True, text=True)
    except subprocess.CalledProcessError as e:
        if e.stderr:
            print(e.stderr, file=sys.stderr, end="")
        die(f"grok trace failed (exit {e.returncode})")

    lines = [ln.strip() for ln in proc.stdout.splitlines() if ln.strip()]
    if not lines:
        if proc.stderr:
            print(proc.stderr, file=sys.stderr, end="")
        die("grok trace produced empty stdout")

    payload = None
    for line in reversed(lines):
        try:
            payload = json.loads(line)
            break
        except json.JSONDecodeError:
            continue
    if payload is None:
        die(f"could not parse json from: {lines[-1]!r}")

    local_path = payload.get("local_path")
    if not local_path:
        die(f"json missing local_path: {payload!r}")
    return Path(local_path)


remote_path = run_trace(args.from_host)

if args.from_host:
    dest = Path(gettempdir()) / remote_path.name
    print(f"rsync {args.from_host}:{remote_path} → {dest}", file=sys.stderr)
    try:
        subprocess.run(
            [
                "rsync",
                "-avP",
                *rsync_ssh(),
                f"{args.from_host}:{remote_path}",
                str(dest),
            ],
            check=True,
        )
    except subprocess.CalledProcessError as e:
        die(f"rsync failed (exit {e.returncode})")
    print(dest)
else:
    print(remote_path)
