#!/usr/bin/env python3
"""Serve your own AlphaFold 3 to Talindrew as a custom fold model.

Talindrew cannot run AlphaFold 3 for you: Google DeepMind licenses the model
parameters for non-commercial use only. If you (or your institution) hold that
licence and have AlphaFold 3 installed, this one-file adapter puts it behind
the contract Talindrew's fold stage speaks, so you can add it under
AI workspace -> Fold -> Use your own model.

The contract (https://www.talindrew.com/guides/custom-models):

    POST /fold            {"sequence": "MKT...", "params": {...}}
      -> 202, header nvcf-reqid: <job id>         (AlphaFold 3 takes minutes)
    GET  /status/<job id>
      -> 202 while running
      -> 200 {"pdb": "...", "confidence": 0.0-1.0, "per_residue_confidence": [...],
              "plddt_scale": "0-1", "ptm": ..., "model": "alphafold3"}
      -> 500 {"detail": "..."} if AlphaFold 3 failed

Requirements: Python 3.10+, Biopython (`pip install biopython`) to turn
AlphaFold 3's mmCIF into PDB, and a working AlphaFold 3 install.

Configuration (environment variables):

    AF3_COMMAND     How to run one prediction. {json} {out} are filled in.
                    Default: python run_alphafold.py --json_path={json} --output_dir={out}
                    Add your --model_dir / --db_dir flags here, e.g.
                    "python /app/alphafold/run_alphafold.py --json_path={json} --output_dir={out}
                     --model_dir=/models --db_dir=/databases"
    ADAPTER_TOKEN   If set, requests must carry "Authorization: Bearer <token>".
                    Put the same token in Talindrew's token field. Strongly advised.
    PORT            Default 8787.
    WORK_DIR        Where job inputs and outputs go. Default: a temp directory.

Talindrew only calls https endpoints, so put this behind TLS (a reverse proxy,
a tunnel such as cloudflared or tailscale funnel, or your cloud's load
balancer). Jobs run one at a time; AlphaFold 3 wants the whole GPU anyway.

Licence of this file: MIT. AlphaFold 3 itself is under Google DeepMind's terms.
"""

from __future__ import annotations

import io
import json
import os
import queue
import re
import shlex
import subprocess
import tempfile
import threading
import uuid
from glob import glob
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
from pathlib import Path

AF3_COMMAND = os.environ.get("AF3_COMMAND", "python run_alphafold.py --json_path={json} --output_dir={out}")
TOKEN = os.environ.get("ADAPTER_TOKEN", "")
PORT = int(os.environ.get("PORT", "8787"))
WORK = Path(os.environ.get("WORK_DIR") or tempfile.mkdtemp(prefix="af3-adapter-"))
AMINO = re.compile(r"^[ACDEFGHIKLMNPQRSTVWY]+$")
MAX_LENGTH = int(os.environ.get("MAX_LENGTH", "2500"))

JOBS: dict[str, dict] = {}
QUEUE: "queue.Queue[str]" = queue.Queue()


def af3_input(job_id: str, sequence: str, params: dict) -> dict:
    """AlphaFold 3's own JSON input format ("alphafold3" dialect, version 1)."""
    seeds = params.get("seeds") or [int(params.get("seed", 1))]
    return {
        "name": f"talindrew_{job_id}",
        "modelSeeds": [int(s) for s in seeds][:5],
        "sequences": [{"protein": {"id": "A", "sequence": sequence}}],
        "dialect": "alphafold3",
        "version": 1,
    }


def cif_to_result(cif_path: str) -> dict:
    """mmCIF (pLDDT 0-100 in B-factors) -> the fold stage's answer (confidence 0-1)."""
    from Bio.PDB import MMCIFParser, PDBIO

    structure = MMCIFParser(QUIET=True).get_structure("af3", cif_path)
    per_residue = []
    for residue in structure[0].get_residues():
        atoms = list(residue.get_atoms())
        if atoms:
            ca = residue["CA"] if "CA" in residue else atoms[0]
            per_residue.append(round(ca.get_bfactor() / 100.0, 4))
    out = io.StringIO()
    writer = PDBIO()
    writer.set_structure(structure)
    writer.save(out)
    result = {
        "pdb": out.getvalue(),
        "confidence": round(sum(per_residue) / len(per_residue), 4) if per_residue else 0.0,
        "per_residue_confidence": per_residue,
        "plddt_scale": "0-1",
        "residues": len(per_residue),
        "model": "alphafold3",
    }
    summary = glob(str(Path(cif_path).parent / "*summary_confidences.json"))
    if summary:
        with open(summary[0]) as fh:
            s = json.load(fh)
        for key in ("ptm", "iptm", "ranking_score"):
            if s.get(key) is not None:
                result[key] = s[key]
    return result


def run_job(job_id: str) -> None:
    job = JOBS[job_id]
    job_dir = WORK / job_id
    out_dir = job_dir / "out"
    out_dir.mkdir(parents=True, exist_ok=True)
    json_path = job_dir / "input.json"
    json_path.write_text(json.dumps(af3_input(job_id, job["sequence"], job["params"])))
    cmd = [part.format(json=json_path, out=out_dir) for part in shlex.split(AF3_COMMAND)]
    job["status"] = "running"
    try:
        run = subprocess.run(cmd, capture_output=True, text=True, timeout=int(os.environ.get("AF3_TIMEOUT", "7200")))
        # AlphaFold 3 writes <name>/<name>_model.cif; the top-ranked model is the one without a seed suffix.
        models = sorted(glob(str(out_dir / "**" / "*_model.cif"), recursive=True), key=len)
        if run.returncode != 0 or not models:
            tail = (run.stderr or run.stdout or "").strip().splitlines()[-5:]
            raise RuntimeError("AlphaFold 3 did not produce a model: " + " | ".join(tail))
        job["result"] = cif_to_result(models[0])
        job["status"] = "done"
    except Exception as exc:  # noqa: BLE001 — reported to the caller, not raised
        job["status"], job["error"] = "failed", str(exc)[:1000]


def worker() -> None:
    while True:
        run_job(QUEUE.get())


class Handler(BaseHTTPRequestHandler):
    def _answer(self, code: int, body: dict | None = None, headers: dict | None = None) -> None:
        self.send_response(code)
        self.send_header("Content-Type", "application/json")
        for k, v in (headers or {}).items():
            self.send_header(k, v)
        self.end_headers()
        if body is not None:
            self.wfile.write(json.dumps(body).encode())

    def _authorised(self) -> bool:
        if not TOKEN:
            return True
        if self.headers.get("Authorization", "") == f"Bearer {TOKEN}":
            return True
        self._answer(401, {"detail": "Missing or wrong bearer token"})
        return False

    def do_POST(self) -> None:  # noqa: N802
        if not self._authorised():
            return
        if self.path.rstrip("/") != "/fold":
            return self._answer(404, {"detail": "POST /fold"})
        try:
            body = json.loads(self.rfile.read(int(self.headers.get("Content-Length", 0))) or b"{}")
        except ValueError:
            return self._answer(400, {"detail": "Body must be JSON"})
        sequence = str(body.get("sequence", "")).strip().upper()
        if not AMINO.match(sequence) or len(sequence) > MAX_LENGTH:
            return self._answer(422, {"detail": f"'sequence' must be 1-{MAX_LENGTH} standard amino acids"})
        job_id = uuid.uuid4().hex
        JOBS[job_id] = {"sequence": sequence, "params": body.get("params") or {}, "status": "queued"}
        QUEUE.put(job_id)
        self._answer(202, {"status": "queued"}, {"nvcf-reqid": job_id})

    def do_GET(self) -> None:  # noqa: N802
        if self.path == "/health":
            return self._answer(200, {"ok": True, "queued": QUEUE.qsize()})
        if not self._authorised():
            return
        match = re.fullmatch(r"/status/([0-9a-f]{32})", self.path)
        job = JOBS.get(match.group(1)) if match else None
        if job is None:
            return self._answer(404, {"detail": "No such job"})
        if job["status"] in ("queued", "running"):
            return self._answer(202, {"status": job["status"]})
        if job["status"] == "failed":
            return self._answer(500, {"detail": job["error"]})
        return self._answer(200, job["result"])

    def log_message(self, fmt: str, *args) -> None:  # quieter than the default
        print(f"[af3-adapter] {self.address_string()} {fmt % args}")


def main() -> None:
    threading.Thread(target=worker, daemon=True).start()
    print(f"AlphaFold 3 adapter on :{PORT} (token {'required' if TOKEN else 'NOT set'}), work dir {WORK}")
    ThreadingHTTPServer(("0.0.0.0", PORT), Handler).serve_forever()


if __name__ == "__main__":
    main()
