Spaces:
Running
Running
Download bioai-platform/backend/app/tools/structure_prep.py from Samad14/bio-nexus-api: direct link, hf CLI and curl.
- Browser
- Download file 15.5 kB
-
https://huggingface.co/spaces/Samad14/bio-nexus-api/resolve/main/bioai-platform/backend/app/tools/structure_prep.py
- Command line
-
hf download hf://spaces/Samad14/bio-nexus-api/bioai-platform/backend/app/tools/structure_prep.py
-
curl -L -o structure_prep.py https://huggingface.co/spaces/Samad14/bio-nexus-api/resolve/main/bioai-platform/backend/app/tools/structure_prep.py
15.5 kB
| """Structure preparation pipeline tools. | |
| Pipeline: fetch β broken chain detection β SWISS-MODEL repair β cleanup β fpocket β CASTp | |
| """ | |
| import asyncio | |
| import io | |
| import logging | |
| import re | |
| import shutil | |
| import subprocess | |
| import tempfile | |
| from dataclasses import dataclass, field | |
| from pathlib import Path | |
| from typing import Any | |
| import httpx | |
| from app.services.identifier_resolution import UNIPROT_RE | |
| from app.services.ssrf import validate_url | |
| logger = logging.getLogger(__name__) | |
| # fpocket remains a compiled binary (installed in the API Dockerfiles). | |
| FPOCKET_BIN = shutil.which("fpocket") or "/usr/local/bin/fpocket" | |
| CASTPFOLD_BASE = "https://cfold.bme.uic.edu/castpfold" | |
| SMR_REPO = "https://swissmodel.expasy.org/repository" | |
| ESMFOLD_API = "https://api-inference.huggingface.co/models/facebook/esmfold_v1" | |
| RCSB_DOWNLOAD = "https://files.rcsb.org/download" | |
| # Input format validation (A4) β reject before any network call. | |
| PDB_ID_RE = re.compile(r"^[A-Za-z0-9]{4}$") | |
| TEMPLATE_RE = re.compile(r"^[A-Za-z0-9][A-Za-z0-9_-]*$") | |
| def validate_pdb_id(pdb_id: str, param_name: str = "pdb_id") -> str: | |
| pdb_id = (pdb_id or "").strip() | |
| if not PDB_ID_RE.match(pdb_id): | |
| raise ValueError(f"{param_name}: expected a 4-character alphanumeric PDB ID, got {pdb_id!r}") | |
| return pdb_id.upper() | |
| def validate_template(template: str, param_name: str = "template") -> str: | |
| template = (template or "").strip() | |
| if not TEMPLATE_RE.match(template): | |
| raise ValueError(f"{param_name}: invalid template ID {template!r}") | |
| return template | |
| # ββ Step 1: Broken chain detection βββββββββββββββββββββββββββββββββββββββββββ | |
| class ChainHealth: | |
| has_missing_residues: bool = False | |
| missing_residue_count: int = 0 | |
| missing_ranges: list[str] = field(default_factory=list) | |
| has_chain_breaks: bool = False | |
| chain_break_count: int = 0 | |
| chain_breaks: list[dict] = field(default_factory=list) | |
| is_broken: bool = False | |
| chains: list[str] = field(default_factory=list) | |
| total_residues: int = 0 | |
| def detect_chain_health(pdb_text: str) -> ChainHealth: | |
| """Detect broken chains: missing residues (REMARK 465) + CA-CA distance gaps.""" | |
| health = ChainHealth() | |
| # --- Parse REMARK 465 (missing residues) --- | |
| remark_lines = [ | |
| line for line in pdb_text.splitlines() | |
| if line.startswith("REMARK 465") | |
| ] | |
| missing_pattern = re.compile( | |
| r"REMARK 465\s+(\S+)\s+(\S+)\s+(\d+)([A-Z]?)\s+(\d+)\s*([A-Z]?)" | |
| ) | |
| for line in remark_lines: | |
| m = missing_pattern.match(line) | |
| if m: | |
| resname = m.group(1) | |
| chain_id = m.group(3) or " " | |
| resnum = int(m.group(5)) | |
| health.missing_residue_count += 1 | |
| health.missing_ranges.append( | |
| f"{chain_id.strip() or ' '}{resnum}{resname}" | |
| ) | |
| health.has_missing_residues = health.missing_residue_count > 0 | |
| # --- CA-CA distance check --- | |
| from Bio.PDB import PDBParser | |
| parser = PDBParser(QUIET=True) | |
| structure = parser.get_structure("pdb", io.StringIO(pdb_text)) | |
| model = structure[0] | |
| health.chains = [c.id for c in model] | |
| for chain in model: | |
| ca_atoms = [ | |
| res["CA"] | |
| for res in chain | |
| if res.id[0] == " " and "CA" in res | |
| ] | |
| ca_atoms.sort(key=lambda a: a.parent.id[1]) | |
| health.total_residues += len(ca_atoms) | |
| for i in range(1, len(ca_atoms)): | |
| prev = ca_atoms[i - 1] | |
| curr = ca_atoms[i] | |
| dist = (prev.coord - curr.coord).tolist() | |
| dist_val = (dist[0] ** 2 + dist[1] ** 2 + dist[2] ** 2) ** 0.5 | |
| if dist_val > 4.2: | |
| health.has_chain_breaks = True | |
| health.chain_break_count += 1 | |
| health.chain_breaks.append({ | |
| "chain": chain.id, | |
| "from_resnum": ca_atoms[i - 1].parent.id[1], | |
| "to_resnum": ca_atoms[i].parent.id[1], | |
| "distance": round(dist_val, 2), | |
| }) | |
| health.is_broken = health.has_missing_residues or health.has_chain_breaks | |
| return health | |
| # ββ Step 2: SWISS-MODEL repair βββββββββββββββββββββββββββββββββββββββββββββββ | |
| def validate_uniprot_accession(accession: str) -> str: | |
| acc = (accession or "").strip().upper() | |
| if not UNIPROT_RE.match(acc): | |
| raise ValueError(f"uniprot_accession: invalid UniProt accession {accession!r}") | |
| return acc | |
| async def swissmodel_fetch_structures(accession: str) -> dict: | |
| """Fetch available structures from SMR Repository for a UniProt accession.""" | |
| acc = validate_uniprot_accession(accession) | |
| url = f"{SMR_REPO}/uniprot/{acc}.json" | |
| validate_url(url) | |
| async with httpx.AsyncClient(timeout=30) as client: | |
| resp = await client.get(url) | |
| if resp.status_code == 404: | |
| return {"models": [], "experimental": []} | |
| resp.raise_for_status() | |
| data = resp.json() | |
| result = data.get("result", {}) | |
| structures = result.get("structures", []) | |
| models = [] | |
| experimental = [] | |
| for s in structures: | |
| entry = { | |
| "template": s.get("template"), | |
| "method": s.get("method"), | |
| "coverage": s.get("coverage"), | |
| "coordinates_url": s.get("coordinates"), | |
| } | |
| if s.get("provider") == "PDB": | |
| experimental.append(entry) | |
| else: | |
| models.append(entry) | |
| return {"models": models, "experimental": experimental, "sequence": result.get("sequence", "")} | |
| async def swissmodel_fetch_pdb(template: str) -> str | None: | |
| """Fetch PDB coordinates from SMR for a template ID.""" | |
| template = validate_template(template, "swissmodel_template") | |
| url = f"{SMR_REPO}/templates/{template}.pdb" | |
| validate_url(url) | |
| try: | |
| async with httpx.AsyncClient(timeout=30) as client: | |
| resp = await client.get(url) | |
| if resp.status_code == 200 and len(resp.text) > 50: | |
| return resp.text | |
| except Exception: | |
| pass | |
| return None | |
| # ββ Step 3: Structure cleanup (pymol2 wheel, Biopython fallback) ββββββββββββ | |
| def pymol_cleanup(pdb_text: str) -> str: | |
| """Remove waters and hetero atoms using the pymol2 Python wheel. | |
| No PyMOL binary or X server required. Falls back to Biopython stripping | |
| (logged loudly, never silently) if pymol2 is unavailable. | |
| """ | |
| try: | |
| return _pymol_cleanup_pymol2(pdb_text) | |
| except Exception as e: | |
| logger.warning("pymol2 cleanup unavailable (%s); using Biopython fallback", e) | |
| return _biopython_cleanup(pdb_text) | |
| def _pymol_cleanup_pymol2(pdb_text: str) -> str: | |
| """Use the importable open-source PyMOL (pymol-open-source-whl) for cleanup.""" | |
| import pymol2 | |
| with tempfile.TemporaryDirectory() as tmpdir: | |
| in_path = Path(tmpdir) / "input.pdb" | |
| out_path = Path(tmpdir) / "clean.pdb" | |
| in_path.write_text(pdb_text) | |
| with pymol2.PyMOL() as p: | |
| p.cmd.load(str(in_path), "struct") | |
| p.cmd.remove("resn HOH") | |
| p.cmd.remove("hetatm") | |
| p.cmd.save(str(out_path), "struct") | |
| result = out_path.read_text() | |
| if len(result) <= 100: | |
| raise RuntimeError("pymol2 produced empty output") | |
| return result | |
| def _biopython_cleanup(pdb_text: str) -> str: | |
| """Strip waters/hetero atoms using Biopython (fallback).""" | |
| from Bio.PDB import PDBParser, PDBIO, Select | |
| class ProteinSelect(Select): | |
| def accept_residue(self, res): | |
| return res.id[0] == " " | |
| parser = PDBParser(QUIET=True) | |
| structure = parser.get_structure("pdb", io.StringIO(pdb_text)) | |
| io_buf = io.BytesIO() | |
| pdb_io = PDBIO() | |
| pdb_io.set_structure(structure) | |
| pdb_io.save(io_buf, ProteinSelect()) | |
| return io_buf.getvalue().decode("utf-8") | |
| # ββ Step 4: fpocket (local binary) βββββββββββββββββββββββββββββββββββββββββββ | |
| class FpocketResult: | |
| pocket_count: int = 0 | |
| pockets: list[dict] = field(default_factory=list) | |
| raw_output: str = "" | |
| status: str = "complete" # complete | unavailable | error | |
| def run_fpocket(pdb_text: str, probe_radius: float = 1.4) -> FpocketResult: | |
| """Run fpocket on PDB text. Returns pocket data.""" | |
| if not Path(FPOCKET_BIN).exists(): | |
| logger.warning("fpocket binary not found at %s β was it installed in the image?", FPOCKET_BIN) | |
| return FpocketResult(raw_output="fpocket not installed", status="unavailable") | |
| with tempfile.TemporaryDirectory() as tmpdir: | |
| in_path = Path(tmpdir) / "input.pdb" | |
| in_path.write_text(pdb_text) | |
| try: | |
| result = subprocess.run( | |
| [FPOCKET_BIN, "-f", str(in_path), "-r", str(probe_radius)], | |
| capture_output=True, | |
| text=True, | |
| timeout=60, | |
| ) | |
| fpocket_out = Path(tmpdir) / "input_out" | |
| return _parse_fpocket_output(fpocket_out, result.stdout + result.stderr) | |
| except subprocess.TimeoutExpired: | |
| return FpocketResult(raw_output="fpocket timed out", status="error") | |
| except Exception as e: | |
| return FpocketResult(raw_output=f"fpocket error: {e}", status="error") | |
| def _parse_fpocket_output(out_dir: Path, raw_output: str) -> FpocketResult: | |
| """Parse fpocket output directory for pocket information.""" | |
| result = FpocketResult(raw_output=raw_output) | |
| info_file = out_dir / "info" / "infos.txt" | |
| if not info_file.exists(): | |
| return result | |
| try: | |
| text = info_file.read_text() | |
| pockets = [] | |
| current_pocket: dict[str, Any] = {} | |
| for line in text.splitlines(): | |
| line = line.strip() | |
| if line.startswith("Pocket"): | |
| if current_pocket: | |
| pockets.append(current_pocket) | |
| pocket_id_match = re.search(r"Pocket\s+(\d+)", line) | |
| current_pocket = { | |
| "id": int(pocket_id_match.group(1)) if pocket_id_match else len(pockets) + 1, | |
| "druggability_score": 0.0, | |
| "volume": 0.0, | |
| "area": 0.0, | |
| "score": 0.0, | |
| "num_residues": 0, | |
| } | |
| elif "Druggability Score" in line: | |
| m = re.search(r":\s*([\d.]+)", line) | |
| if m: | |
| current_pocket["druggability_score"] = float(m.group(1)) | |
| elif "Volume" in line: | |
| m = re.search(r":\s*([\d.]+)", line) | |
| if m: | |
| current_pocket["volume"] = float(m.group(1)) | |
| elif "Area" in line: | |
| m = re.search(r":\s*([\d.]+)", line) | |
| if m: | |
| current_pocket["area"] = float(m.group(1)) | |
| elif "Score" in line and "Drug" not in line: | |
| m = re.search(r":\s*([\d.]+)", line) | |
| if m: | |
| current_pocket["score"] = float(m.group(1)) | |
| elif "Number of residues" in line: | |
| m = re.search(r":\s*(\d+)", line) | |
| if m: | |
| current_pocket["num_residues"] = int(m.group(1)) | |
| if current_pocket: | |
| pockets.append(current_pocket) | |
| result.pockets = pockets | |
| result.pocket_count = len(pockets) | |
| except Exception as e: | |
| logger.warning("Failed to parse fpocket output: %s", e) | |
| return result | |
| # ββ Step 5: CASTp (remote async via CASTpFold) ββββββββββββββββββββββββββββββ | |
| async def castp_submit(pdb_text: str, probe_radius: float = 1.4) -> dict: | |
| """Submit PDB to CASTpFold server for pocket analysis.""" | |
| url = f"{CASTPFOLD_BASE}/compute" | |
| validate_url(url) | |
| async with httpx.AsyncClient(timeout=60) as client: | |
| files = {"pdb_file": ("structure.pdb", pdb_text.encode(), "text/plain")} | |
| data = {"radius": str(probe_radius)} | |
| resp = await client.post(url, files=files, data=data) | |
| resp.raise_for_status() | |
| text = resp.text | |
| job_match = re.search(r"result/([a-f0-9\-]+)", text) or re.search( | |
| r'job[_\-]?id["\s:=]+["\']?([a-f0-9\-]+)', text | |
| ) | |
| if job_match: | |
| return {"job_id": job_match.group(1), "status": "submitted"} | |
| return {"status": "complete", "raw_html": text, "job_id": None} | |
| async def castp_poll(job_id: str) -> dict: | |
| """Poll CASTpFold for job results.""" | |
| if not re.fullmatch(r"[a-f0-9\-]+", job_id or ""): | |
| raise ValueError(f"castp job_id: invalid format {job_id!r}") | |
| url = f"{CASTPFOLD_BASE}/result/{job_id}" | |
| validate_url(url) | |
| async with httpx.AsyncClient(timeout=30) as client: | |
| resp = await client.get(url) | |
| resp.raise_for_status() | |
| text = resp.text | |
| pockets = _parse_castp_html(text) | |
| if pockets: | |
| return {"status": "complete", "pockets": pockets} | |
| return {"status": "running"} | |
| def _parse_castp_html(html: str) -> list[dict]: | |
| """Parse pocket data from CASTpFold result HTML.""" | |
| pockets = [] | |
| row_pattern = re.compile( | |
| r"<tr[^>]*>.*?<td[^>]*>\s*(\d+)\s*</td>" | |
| r".*?<td[^>]*>\s*([\d.]+)\s*</td>" | |
| r".*?<td[^>]*>\s*([\d.]+)\s*</td>.*?</tr>", | |
| re.DOTALL | re.IGNORECASE, | |
| ) | |
| for m in row_pattern.finditer(html): | |
| pockets.append({ | |
| "id": int(m.group(1)), | |
| "area_sa": float(m.group(2)), | |
| "volume_sa": float(m.group(3)), | |
| }) | |
| return pockets | |
| # ββ Pipeline orchestrator ββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| async def fetch_pdb_text(pdb_id: str) -> str: | |
| """Fetch PDB from RCSB.""" | |
| pdb_id = validate_pdb_id(pdb_id) | |
| url = f"{RCSB_DOWNLOAD}/{pdb_id}.pdb" | |
| validate_url(url) | |
| async with httpx.AsyncClient(timeout=30) as client: | |
| resp = await client.get(url) | |
| resp.raise_for_status() | |
| return resp.text | |
| async def esmfold_predict(sequence: str) -> str | None: | |
| """Predict structure from amino acid sequence using ESMFold via HF Inference API.""" | |
| import asyncio | |
| import os | |
| validate_url(ESMFOLD_API) | |
| hf_token = os.environ.get("HF_TOKEN") or os.environ.get("HUGGING_FACE_HUB_TOKEN") | |
| headers = {} | |
| if hf_token: | |
| headers["Authorization"] = f"Bearer {hf_token}" | |
| async with httpx.AsyncClient(timeout=300) as client: | |
| resp = await client.post( | |
| ESMFOLD_API, | |
| json={"inputs": sequence}, | |
| headers=headers, | |
| ) | |
| if resp.status_code == 503: | |
| data = resp.json() | |
| wait_time = min(data.get("estimated_time", 30), 120) | |
| await asyncio.sleep(wait_time) | |
| resp = await client.post( | |
| ESMFOLD_API, | |
| json={"inputs": sequence}, | |
| headers=headers, | |
| ) | |
| resp.raise_for_status() | |
| data = resp.json() | |
| pdb_text = data.get("pdb", "") if isinstance(data, dict) else "" | |
| if pdb_text and len(pdb_text) > 50: | |
| return pdb_text | |
| return None | |