#!/usr/bin/env python3
"""Command-line client for the ArmandoMD REST API. Python 3.9+, standard library only.

Configuration (environment):
    ARMANDOMD_URL    the site, e.g. https://gqc.quimica.unlp.edu.ar/armandomd/
    ARMANDOMD_TOKEN  a personal token (create it in the site, under your email)

Commands:
    me                                  who you are and your quotas
    get PATH                            GET /api/v1/PATH and print the JSON
    post PATH [JSON]                    POST a JSON body (or "-" to read it from stdin)
    wait JOB_ID                         wait until a job finishes
    download EXPORT_ID FILE             download an export package
    pipeline CONFIG.json [--out FILE]   PDB code -> receptor -> ligands -> docking -> parameters
                                        -> system -> MD inputs -> export, waiting for each step
"""

from __future__ import annotations

import argparse
import json
import os
import sys
import time
import urllib.error
import urllib.request

POLL_SECONDS = 5
DONE = ("completed", "failed", "cancelled")


class ApiError(Exception):
    def __init__(self, status, body):
        self.status, self.body = status, body
        message = body.get("error", "") if isinstance(body, dict) else str(body)
        if isinstance(body, dict) and body.get("errors"):
            message += " " + json.dumps(body["errors"], ensure_ascii=False)
        super().__init__(f"HTTP {status}: {message}")


class Client:
    def __init__(self, url: str, token: str, log=None):
        self.base = url.rstrip("/") + "/api/v1/"
        self.token = token
        self.log = log or (lambda message: print(message, file=sys.stderr))

    def _request(self, method: str, path: str, body=None, raw=False):
        data = json.dumps(body).encode() if body is not None else None
        request = urllib.request.Request(self.base + path.lstrip("/"), data=data, method=method)
        request.add_header("Authorization", f"Bearer {self.token}")
        request.add_header("Accept", "application/json")
        if data is not None:
            request.add_header("Content-Type", "application/json")
        try:
            with urllib.request.urlopen(request, timeout=300) as response:  # noqa: S310 - the user's own site
                content = response.read()
        except urllib.error.HTTPError as exc:
            content = exc.read()
            try:
                parsed = json.loads(content)
            except ValueError:
                parsed = content.decode(errors="replace")[:500]
            raise ApiError(exc.code, parsed) from None
        return content if raw else json.loads(content)

    def get(self, path: str):
        return self._request("GET", path)

    def post(self, path: str, body: dict | None = None):
        return self._request("POST", path, body or {})

    def download(self, path: str, target: str) -> int:
        content = self._request("GET", path, raw=True)
        with open(target, "wb") as fh:
            fh.write(content)
        return len(content)

    def wait(self, job_id: str, timeout: float = 7200) -> dict:
        """Poll a job until it finishes; raise if it did not complete."""
        deadline = time.monotonic() + timeout
        last = None
        while True:
            job = self.get(f"jobs/{job_id}/")
            if job["status"] != last:
                self.log(f"  {job['name']}: {job['status']}")
                last = job["status"]
            if job["status"] in DONE:
                if job["status"] != "completed":
                    messages = "; ".join(s["message"] for s in job["stages"] if s["message"])
                    raise RuntimeError(f"{job['name']}: {job['status']} {messages}")
                return job
            if time.monotonic() > deadline:
                raise TimeoutError(f"{job['name']} did not finish in {timeout:.0f} s")
            time.sleep(POLL_SECONDS)

    def pipeline(self, config: dict, out: str | None = None) -> dict:
        """The whole flow of the web interface, step by step. Returns the ids of what it made."""
        made = {}
        self.log(f"structure {config['pdb_id']}")
        existing = self.get(f"structures/?pdb_id={config['pdb_id']}")["structures"]
        structure = existing[0] if existing else self.post("structures/", {"pdb_id": config["pdb_id"]})
        made["structure"] = structure["id"]

        self.log("receptor preparation")
        job = self.wait(self.post(f"structures/{structure['id']}/prepare/", config.get("preparation", {}))["job"]["id"])
        made["receptor"] = receptor_id = next(s["result"]["receptor_id"] for s in job["stages"]
                                              if "receptor_id" in s["result"])  # fmt: skip
        box_body = config.get("box") or {"from_reference_ligand": True}
        made["box"] = box_id = self.post(f"receptors/{receptor_id}/boxes/", box_body)["id"]

        self.log("ligands")
        ligand_set = self.post("ligand-sets/", config["ligands"])
        self.wait(ligand_set["job"]["id"])
        made["ligand_set"] = ligand_set["id"]

        self.log("docking")
        docking_body = {"name": config.get("name", config["pdb_id"]), "receptor": receptor_id, "box": box_id,
                        "ligand_set": ligand_set["id"], **config.get("docking", {})}  # fmt: skip
        docking = self.post("dockings/", docking_body)
        self.wait(docking["job"]["id"])
        made["docking"] = docking["id"]
        poses = [p for r in self.get(f"dockings/{docking['id']}/")["results"] for p in r["poses"]]
        if not poses:
            raise RuntimeError("the docking produced no poses")
        best = min(poses, key=lambda p: p["affinity"])
        self.log(f"  best pose {best['affinity']} kcal/mol, RMSD to the crystal {best['rmsd_reference']}")
        self.post(f"dockings/{docking['id']}/select/", {"poses": [best["id"]]})
        made["pose"] = best["id"]

        self.log("ligand parameters (AM1-BCC)")
        parametrized = self.post(f"dockings/{docking['id']}/parametrize/", config.get("parameters", {}))
        if parametrized.get("job"):
            self.wait(parametrized["job"]["id"])
        made["parameters"] = parameters_id = next(iter(parametrized["parameters"].values()))

        self.log("MD system")
        system = self.post("systems/", {"name": config.get("name", config["pdb_id"]), "pose": best["id"],
                                        "parameters": parameters_id, **config.get("system", {})})  # fmt: skip
        self.wait(system["job"]["id"])
        made["system"] = system["id"]

        self.log("MD inputs")
        inputs = self.post("inputs/", {"system": system["id"], **config.get("inputs", {})})
        self.wait(inputs["job"]["id"])
        made["inputs"] = inputs["id"]

        self.log("export")
        export = self.post("exports/", {"inputs": inputs["id"], **config.get("export", {})})
        self.wait(export["job"]["id"])
        made["export"] = export["id"]
        if out:
            size = self.download(f"exports/{export['id']}/download/", out)
            self.log(f"package written to {out} ({size} bytes)")
        return made


def _client() -> Client:
    url, token = os.environ.get("ARMANDOMD_URL"), os.environ.get("ARMANDOMD_TOKEN")
    if not url or not token:
        sys.exit("Set ARMANDOMD_URL and ARMANDOMD_TOKEN (create the token in the site, under your email).")
    return Client(url, token)


def main(argv=None):
    parser = argparse.ArgumentParser(description="ArmandoMD from the command line.")
    commands = parser.add_subparsers(dest="command", required=True)
    commands.add_parser("me")
    commands.add_parser("get").add_argument("path")
    post = commands.add_parser("post")
    post.add_argument("path")
    post.add_argument("body", nargs="?", default="{}")
    commands.add_parser("wait").add_argument("job_id")
    download = commands.add_parser("download")
    download.add_argument("export_id")
    download.add_argument("file")
    pipeline = commands.add_parser("pipeline")
    pipeline.add_argument("config")
    pipeline.add_argument("--out")
    args = parser.parse_args(argv)

    client = _client()
    try:
        if args.command == "me":
            result = client.get("me/")
        elif args.command == "get":
            result = client.get(args.path)
        elif args.command == "post":
            body = sys.stdin.read() if args.body == "-" else args.body
            result = client.post(args.path, json.loads(body))
        elif args.command == "wait":
            result = client.wait(args.job_id)
        elif args.command == "download":
            result = {"bytes": client.download(f"exports/{args.export_id}/download/", args.file)}
        else:
            with open(args.config) as fh:
                result = client.pipeline(json.load(fh), out=args.out)
    except (ApiError, RuntimeError, TimeoutError) as exc:
        sys.exit(str(exc))
    print(json.dumps(result, indent=2, ensure_ascii=False))


if __name__ == "__main__":
    main()
