#!/usr/bin/env python3
"""Publish and restore StarGit Markdown packages. Python 3; no dependencies."""

import argparse, hashlib, json, os, re, shlex, subprocess, sys, getpass, tempfile, warnings
from pathlib import Path, PurePosixPath
from urllib.parse import quote, urlsplit, unquote
from urllib.request import Request, build_opener, HTTPRedirectHandler
from urllib.error import HTTPError


def sha(data):
    return hashlib.sha256(data).hexdigest()


def git(root, *args):
    return (
        subprocess.check_output(
            ["git", "-C", str(root), *args], stderr=subprocess.DEVNULL
        )
        .decode()
        .strip()
    )


def safe_path(value):
    if (
        not isinstance(value, str)
        or not value
        or len(value) > 600
        or "\\" in value
        or ":" in value
        or any(ord(c) < 32 for c in value)
    ):
        raise ValueError("Unsafe document path")
    p = PurePosixPath(value)
    if (
        p.is_absolute()
        or any(x in ("", ".", "..", ".git") for x in value.split("/"))
        or p.suffix.lower() not in (".md", ".markdown")
    ):
        raise ValueError("Unsafe document path: " + value)
    return value


def target_path(root, relative):
    relative = safe_path(relative)
    p = root
    for part in relative.split("/"):
        p = p / part
        if p.is_symlink():
            raise ValueError("Refusing symlink: " + relative)
    if not p.resolve().is_relative_to(root):
        raise ValueError("Path leaves the project")
    return p


def saved_credentials(server):
    identity = sha(server.rstrip("/").encode())[:20]
    return Path.home() / ".config" / "stargit" / "credentials" / (identity + ".json")


def key(args):
    value = os.environ.get("STARGIT_API_KEY", "")
    path = Path(args.credentials_file).expanduser() if args.credentials_file else None
    if not path and not value:
        candidate = saved_credentials(args.server)
        if candidate.is_file():
            path = candidate
    if path:
        value = ""
        if path.suffix == ".json":
            data = json.loads(path.read_text())
            value = data.get("STARGIT_API_KEY") or data.get("api_key", "")
        else:
            for line in path.read_text().splitlines():
                if line.startswith("STARGIT_API_KEY="):
                    parts = shlex.split(line.split("=", 1)[1], comments=True)
                    value = parts[0] if parts else ""
    if not isinstance(value, str) or not value.strip():
        raise ValueError(
            "No StarGit API key configured. Create one at "
            + args.server.rstrip("/")
            + "/api-keys#api-keys, then run: python3 stargit_handover.py --server "
            + shlex.quote(args.server)
            + " auth. You can also use STARGIT_API_KEY or --credentials-file."
        )
    return value.strip()


class NoRedirect(HTTPRedirectHandler):
    def redirect_request(self, *args, **kwargs):
        raise ValueError("Refusing API redirect with credentials")


class API:
    def __init__(self, args, token=None):
        self.base = args.server.rstrip("/")
        url = urlsplit(self.base)
        if (
            url.username
            or url.password
            or url.query
            or url.fragment
            or (
                url.scheme != "https"
                and not (
                    url.scheme == "http" and url.hostname in ("127.0.0.1", "localhost")
                )
            )
        ):
            raise ValueError("Use an HTTPS StarGit server")
        self.token = token if token is not None else key(args)
        self.opener = build_opener(NoRedirect())

    def call(self, path, data=None, request_key=None, raw=False):
        headers = {
            "Authorization": "Bearer " + self.token,
            "Accept": "application/json",
        }
        if data is not None:
            headers["Content-Type"] = "application/json"
        if request_key:
            headers["Idempotency-Key"] = request_key
        req = Request(
            self.base + path,
            data=json.dumps(data).encode() if data is not None else None,
            headers=headers,
        )
        try:
            with self.opener.open(req, timeout=90) as r:
                body = r.read(17_000_001)
        except HTTPError as exc:
            try:
                message = json.loads(exc.read(4096)).get("error", "Request failed")
            except Exception:
                message = "Request failed"
            raise ValueError(f"StarGit HTTP {exc.code}: {message}") from None
        if len(body) > 17_000_000:
            raise ValueError("Response exceeds limit")
        return body if raw else json.loads(body)


def authenticate(args):
    if args.credentials_file:
        raise ValueError(
            "Use auth without --credentials-file to save a new key; existing credential files already work with list and pull."
        )
    if not sys.stdin.isatty():
        raise ValueError(
            "Run auth yourself in an interactive terminal. Paste the key at its hidden prompt, not in agent chat."
        )
    print(
        "Create a StarGit account API key: "
        + args.server.rstrip("/")
        + "/api-keys#api-keys"
    )
    with warnings.catch_warnings():
        warnings.simplefilter("error", getpass.GetPassWarning)
        try:
            token = getpass.getpass("Paste StarGit API key (hidden): ").strip()
        except getpass.GetPassWarning:
            raise ValueError(
                "Hidden terminal input is unavailable; nothing was saved."
            ) from None
    if not token:
        raise ValueError("No key entered; nothing saved")
    api = API(args, token=token)
    identity = api.call("/api/handovers/auth")
    if args.repo:
        api.call("/api/repos/" + quote(args.repo, safe="") + "/handovers")
    path = saved_credentials(api.base)
    if path.parent.is_symlink() or path.is_symlink():
        raise ValueError("Refusing a symlink credential destination")
    path.parent.mkdir(parents=True, exist_ok=True)
    path.parent.chmod(0o700)
    fd, temporary = tempfile.mkstemp(prefix=".key-", dir=path.parent)
    try:
        os.fchmod(fd, 0o600)
        with os.fdopen(fd, "w") as stream:
            json.dump({"api_key": token, "server": api.base}, stream)
        os.replace(temporary, path)
    finally:
        if os.path.exists(temporary):
            os.unlink(temporary)
    print(
        "Connected as "
        + str(identity["account"])
        + ". Key saved privately for this StarGit server."
    )
    if args.repo:
        print(
            "Repository access verified. You can now run the handover preview and restore commands."
        )
    if os.environ.get("STARGIT_API_KEY"):
        print(
            "STARGIT_API_KEY is also set and will take precedence over the saved key."
        )


def collect(root, paths, follow):
    selected = {}
    pending = []
    missing = set()
    for value in paths:
        p = Path(value).expanduser()
        if not p.is_absolute():
            p = root / p
        relative = p.relative_to(root).as_posix()
        target_path(root, relative)
        pending.append(relative)
    roots = set(pending)
    while pending:
        relative = pending.pop(0)
        if relative in selected:
            continue
        path = target_path(root, relative)
        body = path.read_bytes()
        if len(body) > 1_000_000:
            raise ValueError("Document exceeds 1 MB: " + relative)
        text = body.decode("utf-8")
        selected[relative] = text
        if len(selected) > 100:
            raise ValueError("More than 100 documents; narrow the publication")
        if not follow or relative not in roots:
            continue
        # Include direct Markdown links and existing backtick .md references.
        links = re.findall(r"\]\(([^)]+)\)", text)
        references = (
            [(v, True) for v in links]
            + [(v, False) for v in re.findall(r"`([^`\n]+\.md)`", text)]
            + [(v, False) for v in re.findall(r"^([A-Za-z0-9_./-]+\.md)$", text, re.M)]
        )
        for value, explicit in references:
            parsed = urlsplit(value)
            if parsed.scheme or parsed.netloc:
                continue
            v = unquote(parsed.path)
            if not v.lower().endswith((".md", ".markdown")):
                continue
            candidate = Path(v)
            if candidate.is_absolute():
                if not candidate.is_relative_to(root):
                    continue
            else:
                candidate = (path.parent / v).resolve()
                if not candidate.exists():
                    candidate = (root / v).resolve()
            if not candidate.is_relative_to(root):
                if explicit:
                    missing.add(value)
                continue
            name = candidate.relative_to(root).as_posix()
            if candidate.is_file():
                target_path(root, name)
                pending.append(name)
            elif explicit:
                missing.add(value)
    if missing:
        raise ValueError("Missing Markdown references: " + ", ".join(sorted(missing)))
    return selected


def publish(args, api):
    root = Path(args.root).expanduser().resolve()
    root = Path(git(root, "rev-parse", "--show-toplevel"))
    files = collect(root, args.file, args.follow_links)
    first = Path(args.file[0]).expanduser()
    first = first if first.is_absolute() else root / first
    data = dict(
        slug=args.name,
        title=args.title,
        entrypoint=first.relative_to(root).as_posix(),
        parent_id=args.parent,
        source=dict(
            commit=git(root, "rev-parse", "HEAD"),
            branch=git(root, "branch", "--show-current"),
            dirty=bool(git(root, "status", "--porcelain")),
        ),
        files=[dict(path=p, content=t) for p, t in sorted(files.items())],
    )
    if args.dry_run:
        print(
            json.dumps(
                dict(
                    entrypoint=data["entrypoint"],
                    source=data["source"],
                    files=[
                        {"path": p, "sha256": sha(t.encode()), "size": len(t.encode())}
                        for p, t in sorted(files.items())
                    ],
                ),
                indent=2,
            )
        )
        return
    request_key = sha(json.dumps(data, sort_keys=True).encode())
    answer = api.call(
        "/api/repos/" + quote(args.repo, safe="") + "/handovers", data, request_key
    )
    print(
        json.dumps(
            {
                "revision": answer["id"],
                "version": answer["version"],
                "documents": len(files),
                "url": api.base + "/repo/" + args.repo + "/handovers/" + answer["id"],
            },
            indent=2,
        )
    )


def restore(root, files, download, dry_run=False):
    """Validate the whole publication before writes; never replace differing files."""
    root = root.resolve()
    prepared = []
    seen = set()
    total = 0
    for entry in files:
        name = safe_path(entry["path"])
        if name.casefold() in seen:
            raise ValueError("Duplicate/case-colliding destination")
        seen.add(name.casefold())
        path = target_path(root, name)
        content = download(name)
        total += len(content)
        if (
            len(content) > 1_000_000
            or total > 8_000_000
            or sha(content) != entry["sha256"]
            or len(content) != entry["size"]
        ):
            raise ValueError("Integrity check failed: " + name)
        content.decode("utf-8")
        if path.exists() and (not path.is_file() or path.read_bytes() != content):
            raise ValueError("Local file differs; no files changed: " + name)
        prepared.append((path, content))
    if not dry_run:
        for path, content in prepared:
            target_path(root, path.relative_to(root).as_posix())
            if path.exists():
                continue
            path.parent.mkdir(parents=True, exist_ok=True)
            # Exclusive creation: a concurrent local edit is never overwritten.
            with path.open("xb") as stream:
                stream.write(content)
    return [p.relative_to(root).as_posix() for p, _ in prepared]


def pull(args, api):
    root = Path(args.root).expanduser().resolve()
    if Path(git(root, "rev-parse", "--show-toplevel")).resolve() != root:
        raise ValueError("--root must be the project checkout root")
    prefix = (
        "/api/repos/"
        + quote(args.repo, safe="")
        + "/handovers/"
        + quote(args.revision, safe="")
    )
    publication = api.call(prefix)
    manifest = publication["manifest"]
    expected = manifest["source"]["commit"]
    actual = git(root, "rev-parse", "HEAD")
    if actual != expected and not args.allow_different_revision:
        raise ValueError(
            "Code revision differs: expected "
            + expected
            + ", found "
            + actual
            + ". Inspect the handover; use --allow-different-revision to download documents intentionally."
        )
    instruction = lambda f: PurePosixPath(f["path"]).name.upper() in (
        "AGENTS.MD",
        "CLAUDE.MD",
        "COPILOT-INSTRUCTIONS.MD",
    ) or f["path"].startswith((".codex/", ".claude/", ".github/instructions/"))
    skipped = [
        f["path"]
        for f in manifest["files"]
        if instruction(f) and not args.include_instructions
    ]
    files = [f for f in manifest["files"] if f["path"] not in skipped]
    names = restore(
        root,
        files,
        lambda p: api.call(prefix + "/files/" + quote(p, safe="/"), raw=True),
        args.dry_run,
    )
    print(
        json.dumps(
            dict(
                revision=publication["id"],
                source_commit=expected,
                local_commit=actual,
                source_had_local_changes=manifest["source"].get("dirty", False),
                dry_run=args.dry_run,
                documents=names,
                instructions_not_installed=skipped,
            ),
            indent=2,
        )
    )
    if not args.dry_run:
        # Local Git metadata, not tracked project files. Instructions stay inactive.
        gitdir = Path(git(root, "rev-parse", "--absolute-git-dir"))
        folder = gitdir / "stargit-handovers"
        folder.mkdir(exist_ok=True)
        record = folder / (publication["id"] + ".json")
        record.write_text(
            json.dumps(
                dict(server=api.base, repository=args.repo, publication=publication),
                indent=2,
            )
        )
        record.chmod(0o600)


def main():
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--server", default="https://stargit.com")
    parser.add_argument("--credentials-file")
    sub = parser.add_subparsers(dest="command", required=True)
    auth = sub.add_parser("auth", help="Save and verify an API key on this computer")
    auth.add_argument("--repo", help="Also verify access to this repository")
    register = sub.add_parser("register")
    register.add_argument("--name", required=True)
    register.add_argument("--remote", required=True)
    listing = sub.add_parser("list")
    listing.add_argument("--repo", required=True)
    pub = sub.add_parser("publish")
    pub.add_argument("--repo", required=True)
    pub.add_argument("--root", default=".")
    pub.add_argument("--file", action="append", required=True)
    pub.add_argument("--name", required=True)
    pub.add_argument("--title", required=True)
    pub.add_argument("--parent")
    pub.add_argument("--follow-links", action="store_true")
    pub.add_argument("--dry-run", action="store_true")
    puller = sub.add_parser("pull")
    puller.add_argument("--repo", required=True)
    puller.add_argument("--revision", required=True)
    puller.add_argument("--root", default=".")
    puller.add_argument("--allow-different-revision", action="store_true")
    puller.add_argument("--dry-run", action="store_true")
    puller.add_argument(
        "--include-instructions",
        action="store_true",
        help="Explicitly install included agent instruction files after review",
    )
    args = parser.parse_args()
    try:
        if args.command == "auth":
            authenticate(args)
            return 0
        api = None if args.command == "publish" and args.dry_run else API(args)
        if args.command == "publish":
            publish(args, api)
        elif args.command == "pull":
            pull(args, api)
        elif args.command == "register":
            print(
                json.dumps(
                    api.call(
                        "/api/handovers/repositories",
                        dict(name=args.name, remote=args.remote),
                    ),
                    indent=2,
                )
            )
        else:
            print(
                json.dumps(
                    api.call("/api/repos/" + quote(args.repo, safe="") + "/handovers"),
                    indent=2,
                )
            )
    except (ValueError, OSError, subprocess.CalledProcessError) as exc:
        print("Error: " + str(exc), file=sys.stderr)
        return 1
    return 0


if __name__ == "__main__":
    sys.exit(main())
