bio-nexus-api / bioai-platform /backend /app /tools /structure_prep.py
Samad14's picture
feat: techspec additions β€” de novo tier-6 branch, structure export, page capture + final synthesis
4939f56
Raw History Blame Contribute Delete
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 ───────────────────────────────────────────
@dataclass
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) ───────────────────────────────────────────
@dataclass
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