Samad14's picture
fix: NGS fallback to Python aligner when native tools missing; add missing re import in sequence_utils
34b3c94
Raw History Blame Contribute Delete
41.5 kB
"""
NGS (Next-Generation Sequencing) Pipeline — 6-step analysis:
1. Quality Control (FastQC or Python fallback)
2. Trimming/Filter (fastp or Python fallback)
3. Alignment (minimap2 + samtools or Python fallback)
4. Variant Calling (bcftools or Python fallback)
5. Annotation (SnpEff or cross-reference lookup)
6. Visualization (upload BAM/VCF/FASTA for igv.js browser)
"""
from __future__ import annotations
import asyncio
import gzip
import json
import logging
import os
import platform
import random
import re
import shutil
import subprocess
import tempfile
from typing import Any
import httpx
from app.tools.base import BaseTool
logger = logging.getLogger(__name__)
_IS_WINDOWS = os.name == "nt" or platform.system() == "Windows"
BIN_DIR = os.path.join(os.path.dirname(__file__), "..", "bin")
PIPELINE_TIMEOUT = 600
REFERENCE_URLS = {
"sars-cov-2": "https://hgdownload.soe.ucsc.edu/goldenPath/wuhCor1/bigZips/wuhCor1.fa.gz",
"lambda": "https://eutils.ncbi.nlm.nih.gov/entrez/eutils/efetch.fcgi?db=nuccore&id=NC_001416&rettype=fasta&retmode=text",
"ecoli-k12": "https://eutils.ncbi.nlm.nih.gov/entrez/eutils/efetch.fcgi?db=nuccore&id=U00096.3&rettype=fasta&retmode=text",
}
REFERENCE_SIZES = {
"sars-cov-2": 29903,
"lambda": 48502,
"ecoli-k12": 4641652,
}
CACHE_DIR = os.path.join(os.path.dirname(__file__), "..", "data", "references")
# ---------------------------------------------------------------------------
# Tool detection
# ---------------------------------------------------------------------------
def _find_tool(name: str) -> str | None:
exe = f"{name}.exe" if _IS_WINDOWS else name
found = shutil.which(name)
if found:
return found
bundled = os.path.join(BIN_DIR, exe)
if os.path.isfile(bundled):
return bundled
return None
def _tool_available(name: str) -> bool:
path = _find_tool(name)
if not path:
return False
if _IS_WINDOWS:
return os.path.isfile(path)
return os.access(path, os.X_OK)
# ---------------------------------------------------------------------------
# Step 1: Quality Control
# ---------------------------------------------------------------------------
def _run_fastqc(fastq_path: str, out_dir: str) -> dict:
if _tool_available("fastqc"):
try:
subprocess.run(
["fastqc", fastq_path, "-o", out_dir, "-q", "--json"],
capture_output=True, text=True, timeout=120,
)
json_file = os.path.join(out_dir, os.path.basename(fastq_path).replace(".fastq", "_fastqc.json"))
if os.path.exists(json_file):
with open(json_file) as f:
return json.load(f)
except Exception as e:
logger.warning("FastQC failed, using Python fallback: %s", e)
return _python_qc(fastq_path)
def _python_qc(fastq_path: str) -> dict:
total_reads = 0
total_bases = 0
gc_count = 0
at_count = 0
q_scores: list[int] = []
read_lengths: list[int] = []
seen_seqs: dict[str, int] = {}
base_quality_by_pos: dict[int, list[int]] = {}
gc_by_window: list[float] = []
line_no = 0
window_seqs: list[str] = []
WINDOW_SIZE = 50
with open(fastq_path) as f:
for line in f:
line_no += 1
if line_no % 4 == 2:
seq = line.strip()
l = len(seq)
read_lengths.append(l)
total_bases += l
total_reads += 1
gc_count += seq.count("G") + seq.count("C") + seq.count("g") + seq.count("c")
at_count += seq.count("A") + seq.count("T") + seq.count("a") + seq.count("t")
seen_seqs[seq] = seen_seqs.get(seq, 0) + 1
window_seqs.append(seq)
if len(window_seqs) >= WINDOW_SIZE:
gc_in_window = sum(
s.count("G") + s.count("C") + s.count("g") + s.count("c")
for s in window_seqs
)
bases_in_window = sum(len(s) for s in window_seqs)
gc_by_window.append(round(gc_in_window / max(bases_in_window, 1) * 100, 2))
window_seqs = []
elif line_no % 4 == 0:
qual_line = line.strip()
for i, ch in enumerate(qual_line):
q = ord(ch) - 33
q_scores.append(q)
if i not in base_quality_by_pos:
base_quality_by_pos[i] = []
base_quality_by_pos[i].append(q)
if window_seqs:
gc_in_window = sum(s.count("G") + s.count("C") + s.count("g") + s.count("c") for s in window_seqs)
bases_in_window = sum(len(s) for s in window_seqs)
gc_by_window.append(round(gc_in_window / max(bases_in_window, 1) * 100, 2))
if total_reads == 0:
return {"error": "Empty FASTQ file", "total_reads": 0}
mean_q = sum(q_scores) / len(q_scores) if q_scores else 0
q20 = sum(1 for q in q_scores if q >= 20) / len(q_scores) * 100 if q_scores else 0
q30 = sum(1 for q in q_scores if q >= 30) / len(q_scores) * 100 if q_scores else 0
gc_pct = gc_count / (gc_count + at_count) * 100 if (gc_count + at_count) > 0 else 0
# Per-position quality for the chart (sample every N positions)
max_pos = max(base_quality_by_pos.keys()) if base_quality_by_pos else 0
sample_step = max(1, max_pos // 100)
quality_by_position = []
for pos in range(0, max_pos + 1, sample_step):
scores = base_quality_by_pos.get(pos, [])
if scores:
quality_by_position.append({
"position": pos,
"mean": round(sum(scores) / len(scores), 1),
"q10": round(sorted(scores)[len(scores) // 4], 1) if len(scores) >= 4 else 0,
"q90": round(sorted(scores)[len(scores) * 3 // 4], 1) if len(scores) >= 4 else 0,
})
overrepresented = sorted(seen_seqs.items(), key=lambda x: -x[1])[:10]
# Read length distribution
length_dist: dict[int, int] = {}
for rl in read_lengths:
bucket = (rl // 10) * 10
length_dist[bucket] = length_dist.get(bucket, 0) + 1
return {
"tool": "python-qc",
"total_reads": total_reads,
"total_bases": total_bases,
"avg_read_length": round(sum(read_lengths) / len(read_lengths), 1),
"min_read_length": min(read_lengths),
"max_read_length": max(read_lengths),
"gc_percent": round(gc_pct, 2),
"mean_quality": round(mean_q, 2),
"min_quality": min(q_scores),
"max_quality": max(q_scores),
"q20_percent": round(q20, 2),
"q30_percent": round(q30, 2),
"quality_by_position": quality_by_position,
"gc_by_window": gc_by_window,
"read_length_distribution": [{"length": k, "count": v} for k, v in sorted(length_dist.items())],
"overrepresented_sequences": [
{"sequence": s[:50], "count": c, "percent": round(c / total_reads * 100, 2)}
for s, c in overrepresented
],
}
# ---------------------------------------------------------------------------
# Step 2: Trimming/Filtering
# ---------------------------------------------------------------------------
def _run_fastp(fastq_in: str, fastq_out: str, report_dir: str) -> dict:
if _tool_available("fastp"):
report_json = os.path.join(report_dir, "fastp_report.json")
try:
subprocess.run(
["fastp", "-i", fastq_in, "-o", fastq_out,
"--json", report_json, "--thread", "1",
"--qualified_quality_phred", "20",
"--length_required", "50"],
capture_output=True, text=True, timeout=120,
)
if os.path.exists(report_json):
with open(report_json) as f:
return json.load(f)
except Exception as e:
logger.warning("fastp failed, using Python fallback: %s", e)
return _python_trim(fastq_in, fastq_out)
def _python_trim(fastq_in: str, fastq_out: str) -> dict:
kept = 0
discarded = 0
line_no = 0
buf: list[str] = []
before_lengths: list[int] = []
after_lengths: list[int] = []
with open(fastq_in) as fin, open(fastq_out, "w") as out:
for line in fin:
line_no += 1
buf.append(line)
if line_no % 4 == 0:
seq = buf[1].strip()
qual = buf[3].strip()
before_lengths.append(len(seq))
if len(seq) < 50:
discarded += 1
buf.clear()
continue
q_scores = [ord(ch) - 33 for ch in qual]
mean_q = sum(q_scores) / len(q_scores) if q_scores else 0
if mean_q >= 20:
out.writelines(buf)
kept += 1
after_lengths.append(len(seq))
else:
discarded += 1
buf.clear()
avg_before = round(sum(before_lengths) / max(len(before_lengths), 1), 1)
avg_after = round(sum(after_lengths) / max(len(after_lengths), 1), 1)
# Quality distribution before/after
q_before = [ord(ch) - 33 for line in open(fastq_in) for ch in line.strip() if line.strip()]
return {
"tool": "python-trim",
"before_filtering": {"total_reads": kept + discarded, "total_bases": sum(before_lengths), "avg_length": avg_before},
"after_filtering": {"total_reads": kept, "total_bases": sum(after_lengths), "avg_length": avg_after},
"reads_discarded": discarded,
"filtering_result": {"low_quality_reads": discarded},
}
# ---------------------------------------------------------------------------
# Step 3: Alignment
# ---------------------------------------------------------------------------
def _run_alignment(fastq_path: str, ref_path: str, tmpdir: str) -> dict:
sam_path = os.path.join(tmpdir, "aligned.sam")
bam_path = os.path.join(tmpdir, "sorted.bam")
bai_path = os.path.join(tmpdir, "sorted.bam.bai")
mm2 = _find_tool("minimap2")
samtools = _find_tool("samtools")
if not mm2:
raise RuntimeError(
"minimap2 not found on PATH. NGS alignment requires minimap2. "
"Install it: https://github.com/lh3/minimap2"
)
if not samtools:
raise RuntimeError(
"samtools not found on PATH. NGS alignment requires samtools. "
"Install it: https://github.com/samtools/samtools"
)
# minimap2: align reads to reference → SAM
with open(sam_path, "w") as sam_out:
proc = subprocess.run(
[mm2, "-ax", "sr", ref_path, fastq_path],
stdout=sam_out, stderr=subprocess.PIPE, timeout=300,
)
if proc.returncode != 0:
raise RuntimeError(
f"minimap2 failed (exit {proc.returncode}): "
f"{proc.stderr.decode('utf-8', errors='replace')[:500]}"
)
# samtools sort: SAM → sorted BAM
proc_sort = subprocess.run(
["samtools", "sort", "-o", bam_path, sam_path],
capture_output=True, timeout=120,
)
if proc_sort.returncode != 0:
raise RuntimeError(
f"samtools sort failed (exit {proc_sort.returncode}): "
f"{proc_sort.stderr.decode('utf-8', errors='replace')[:500]}"
)
# samtools index: sorted BAM → BAI
proc_idx = subprocess.run(
["samtools", "index", bam_path],
capture_output=True, timeout=60,
)
if proc_idx.returncode != 0:
raise RuntimeError(
f"samtools index failed (exit {proc_idx.returncode}): "
f"{proc_idx.stderr.decode('utf-8', errors='replace')[:500]}"
)
# Verify output files exist and are non-empty
for path, label in [(bam_path, "BAM"), (bai_path, "BAI")]:
if not os.path.exists(path) or os.path.getsize(path) == 0:
raise RuntimeError(f"Native alignment produced empty {label} file: {path}")
stats = _parse_alignment_stats(sam_path)
stats["tool"] = "minimap2+samtools"
stats["bam_path"] = bam_path
stats["bai_path"] = bai_path
stats["sam_path"] = sam_path
stats["read_region"] = _compute_read_region(sam_path)
return stats
def _python_alignment(fastq_path: str, ref_path: str, sam_path: str) -> dict:
"""Naive seed-and-extend aligner. Produces SAM only — no BAM, no BAI.
Used only when native tools are unavailable for diagnostic/debugging purposes.
"""
ref_name = "unknown"
ref_seq_lines = []
with open(ref_path) as rf:
for line in rf:
if line.startswith(">"):
ref_name = line.strip().split()[0][1:].split()[0]
else:
ref_seq_lines.append(line.strip())
ref_seq = "".join(ref_seq_lines)
ref_len = len(ref_seq) if ref_seq else 30000
# Build minimap2-style alignment: match reads against reference
total = 0
mapped = 0
unmapped = 0
reads: list[tuple[str, str, str, int, str, str]] = []
with open(fastq_path) as fin:
line_no = 0
qname = ""
seq = ""
qual = ""
for line in fin:
line_no += 1
if line_no % 4 == 1:
qname = line.strip().lstrip("@")
elif line_no % 4 == 2:
seq = line.strip()
elif line_no % 4 == 0:
qual = line.strip()
total += 1
reads.append((qname, seq, qual))
# Simple seed-and-extend alignment: find best match in reference
import random as _rnd
for idx in range(len(reads)):
qname, seq, qual = reads[idx]
if len(seq) < 20:
unmapped += 1
continue
# Take a 20bp seed from the read and scan the reference
seed = seq[:20].upper()
best_pos = -1
best_score = 0
# Scan reference with a sliding window (sample positions for speed)
step = max(1, ref_len // 500)
for pos in range(0, ref_len - len(seq), step):
ref_window = ref_seq[pos:pos + len(seq)]
matches = sum(1 for a, b in zip(seed, ref_window) if a == b)
if matches > best_score:
best_score = matches
best_pos = pos
# Refine best position
if best_pos >= 0 and best_score >= 10:
# Calculate CIGAR - use M for both match and mismatch (SAM spec)
ref_segment = ref_seq[best_pos:best_pos + len(seq)]
cigar_ops = []
match_count = 0
for i in range(min(len(seq), len(ref_segment))):
if seq[i].upper() == ref_segment[i].upper():
match_count += 1
else:
if match_count > 0:
cigar_ops.append(f"{match_count}M")
match_count = 0
cigar_ops.append("1M")
if match_count > 0:
cigar_ops.append(f"{match_count}M")
cigar = "".join(cigar_ops) if cigar_ops else f"{len(seq)}M"
# Collapse consecutive M operations
prev = None
while prev != cigar:
prev = cigar
cigar = re.sub(r'(\d+)M(\d+)M', lambda m: f"{int(m.group(1)) + int(m.group(2))}M", cigar)
mapq = min(60, best_score * 3)
if len(qual) < len(seq):
qual = qual + "I" * (len(seq) - len(qual))
qual = qual[:len(seq)]
flag = 0
reads[idx] = (qname, seq, qual, best_pos + 1, cigar, flag, mapq)
mapped += 1
else:
unmapped += 1
# Write SAM only
with open(sam_path, "w") as out:
out.write(f"@HD\tVN:1.6\tSO:coordinate\n")
out.write(f"@SQ\tSN:{ref_name}\tLN:{ref_len}\n")
for read in reads:
if len(read) == 3:
qname, seq, qual = read
if len(qual) < len(seq):
qual = qual + "I" * (len(seq) - len(qual))
qual = qual[:len(seq)]
out.write(f"{qname}\t4\t*\t0\t0\t*\t*\t0\t0\t{seq}\t{qual}\n")
else:
qname, seq, qual, pos, cigar, flag, mapq = read
out.write(f"{qname}\t{flag}\t{ref_name}\t{pos}\t{mapq}\t{cigar}\t*\t0\t0\t{seq}\t{qual}\n")
read_region = _compute_read_region(sam_path)
logger.warning(
"DEGRADED MODE: used pure-Python aligner (no minimap2/samtools). "
"SAM output only — no BAM/BAI. Genome browser visualization unavailable."
)
return {
"tool": "python-alignment",
"mapped_reads": mapped,
"unmapped_reads": unmapped,
"total_alignments": total,
"sam_path": sam_path,
"read_region": read_region,
}
def _parse_alignment_stats(sam_path: str) -> dict:
mapped = 0
unmapped = 0
total = 0
with open(sam_path) as f:
for line in f:
if line.startswith("@"):
continue
total += 1
parts = line.strip().split("\t", maxsplit=2)
if len(parts) >= 2:
flag = int(parts[1])
if flag & 4:
unmapped += 1
else:
mapped += 1
return {"mapped_reads": mapped, "unmapped_reads": unmapped, "total_alignments": total}
def _compute_read_region(sam_path: str) -> str:
"""Parse SAM to find min/max mapped positions and return a locus string for igv.js."""
min_pos = 999999999
max_pos = 0
ref_name = ""
with open(sam_path) as f:
for line in f:
if line.startswith("@SQ"):
for p in line.split("\t"):
if p.startswith("SN:"):
ref_name = p[3:]
continue
if line.startswith("@"):
continue
parts = line.strip().split("\t")
if len(parts) < 6:
continue
flag = int(parts[1])
if flag & 4:
continue
rname = parts[2]
pos = int(parts[3])
cigar = parts[5]
ref_span = sum(int(l) for l, op in re.findall(r'(\d+)([MDN=])', cigar))
if ref_span == 0:
ref_span = 100
min_pos = min(min_pos, pos)
max_pos = max(max_pos, pos + ref_span)
if not ref_name:
ref_name = rname
if max_pos > min_pos and ref_name:
pad = max(50, (max_pos - min_pos) // 10)
return f"{ref_name}:{max(1, min_pos - pad)}-{max_pos + pad}"
if ref_name:
return f"{ref_name}:1-500"
return ""
# ---------------------------------------------------------------------------
# Step 4: Variant Calling
# ---------------------------------------------------------------------------
def _run_variant_calling(sam_path: str, ref_path: str, tmpdir: str) -> dict:
vcf_path = os.path.join(tmpdir, "variants.vcf")
bcftools = _find_tool("bcftools")
bam_path = os.path.join(tmpdir, "sorted.bam")
if bcftools and os.path.exists(bam_path):
# bcftools mpileup → bcftools call
proc = subprocess.run(
["bcftools", "mpileup", "-f", ref_path, bam_path, "-O", "u"],
capture_output=True, timeout=120,
)
if proc.returncode != 0:
logger.warning(
"bcftools mpileup returned %d: %s",
proc.returncode,
proc.stderr.decode("utf-8", errors="replace")[:300],
)
else:
call_proc = subprocess.run(
["bcftools", "call", "-mv", "-O", "v"],
input=proc.stdout, capture_output=True, timeout=120,
)
if call_proc.returncode != 0:
logger.warning(
"bcftools call returned %d: %s",
call_proc.returncode,
call_proc.stderr.decode("utf-8", errors="replace")[:300],
)
else:
with open(vcf_path, "w") as f:
f.write(call_proc.stdout.decode("utf-8", errors="replace"))
variants = _parse_vcf(vcf_path)
return {
"tool": "bcftools",
"vcf_path": vcf_path,
"variants": variants,
"total_variants": len(variants),
}
# Fallback: pileup-based variant calling from SAM (text-only, no binary formats)
return _python_variant_calling(sam_path, ref_path, vcf_path)
def _python_variant_calling(sam_path: str, ref_path: str, vcf_path: str) -> dict:
ref_lines = open(ref_path).readlines()
ref_name = "unknown"
ref_seq_lines = []
for line in ref_lines:
if line.startswith(">"):
ref_name = line.strip().split()[0][1:].split()[0]
else:
ref_seq_lines.append(line.strip())
ref = "".join(ref_seq_lines)
pileup: dict[int, dict[str, int]] = {}
depth_by_pos: dict[int, int] = {}
with open(sam_path) as f:
for line in f:
if line.startswith("@"):
continue
parts = line.strip().split("\t")
if len(parts) < 6:
continue
flag = int(parts[1])
if flag & 4:
continue
pos = int(parts[3])
cigar = parts[5]
seq = parts[9]
genome_pos = pos - 1
ops = re.findall(r"(\d+)([MIDNSHPX=])", cigar)
offset = 0
for length, op in ops:
l = int(length)
if op == "M":
for i in range(l):
p = genome_pos + i
if p < len(ref):
base = seq[offset + i].upper() if offset + i < len(seq) else "N"
if p not in pileup:
pileup[p] = {"A": 0, "C": 0, "G": 0, "T": 0}
depth_by_pos[p] = depth_by_pos.get(p, 0) + 1
if base in pileup[p]:
pileup[p][base] += 1
offset += l
elif op == "I":
offset += l
elif op == "S":
offset += l
min_depth = 2
min_alt_freq = 0.2
variants = []
for pos in sorted(pileup.keys()):
counts = pileup[pos]
depth = depth_by_pos.get(pos, sum(counts.values()))
if depth < min_depth:
continue
ref_base = ref[pos].upper() if pos < len(ref) else "N"
total = sum(counts.get(b, 0) for b in "ACGT")
if total == 0:
continue
for base in "ACGT":
if base == ref_base:
continue
alt_count = counts.get(base, 0)
freq = alt_count / total
if freq >= min_alt_freq:
variants.append({
"pos": pos + 1, "ref": ref_base, "alt": base,
"depth": depth, "alt_count": alt_count, "freq": round(freq, 4),
})
variants.sort(key=lambda v: -v["freq"])
variants = variants[:50]
# Write proper VCF with header
with open(vcf_path, "w") as f:
f.write("##fileformat=VCFv4.2\n")
f.write("##source=ngs-pipeline-python\n")
f.write(f"##reference={ref_name}\n")
f.write("#CHROM\tPOS\tID\tREF\tALT\tQUAL\tFILTER\tINFO\n")
for v in variants:
f.write(f"{ref_name}\t{v['pos']}\t.\t{v['ref']}\t{v['alt']}\t.\tPASS\tDP={v['depth']};AF={v['freq']}\n")
return {
"tool": "python-variant-calling",
"vcf_path": vcf_path,
"variants": variants,
"total_variants": len(variants),
}
def _parse_vcf(vcf_path: str) -> list[dict]:
variants = []
with open(vcf_path) as f:
for line in f:
if line.startswith("#"):
continue
parts = line.strip().split("\t")
if len(parts) < 5:
continue
info = {}
for item in parts[7].split(";"):
if "=" in item:
k, v = item.split("=", 1)
info[k] = v
depth = int(info.get("DP", "0"))
af = float(info.get("AF", "0"))
variants.append({
"pos": int(parts[1]),
"ref": parts[3],
"alt": parts[4],
"depth": depth,
"alt_count": round(depth * af) if depth else 0,
"freq": af,
})
return variants
# ---------------------------------------------------------------------------
# Step 5: Annotation
# ---------------------------------------------------------------------------
def _run_annotation(vcf_path: str, ref_name: str, tmpdir: str) -> dict:
annotated_path = os.path.join(tmpdir, "annotated.vcf")
snpeff = _find_tool("snpeff")
if snpeff:
try:
subprocess.run(
[snpeff, "ann", ref_name, vcf_path],
capture_output=True, text=True, timeout=120,
)
except Exception as e:
logger.warning("SnpEff failed, using basic annotation: %s", e)
return _basic_annotation(vcf_path, ref_name)
KNOWN_VARIANTS = {
"sars-cov-2": {
23403: {"gene": "S", "mutation": "D614G", "significance": "Increased transmissibility", "protein_change": "Asp614Gly"},
28881: {"gene": "N", "mutation": "R203K", "significance": "Common variant", "protein_change": "Arg203Lys"},
28882: {"gene": "N", "mutation": "G204R", "significance": "Common variant", "protein_change": "Gly204Arg"},
21563: {"gene": "ORF1ab", "mutation": "P13L", "significance": "Early divergence marker", "protein_change": "Pro13Leu"},
28883: {"gene": "N", "mutation": "G204R", "significance": "Common variant", "protein_change": "Gly204Arg"},
26245: {"gene": "ORF1ab", "mutation": "R203K", "significance": "Common in Wuhan-Hu-1", "protein_change": "Arg203Lys"},
29742: {"gene": "ORF10", "mutation": "G12V", "significance": "Minor variant", "protein_change": "Gly12Val"},
},
"lambda": {},
"ecoli-k12": {},
}
def _basic_annotation(vcf_path: str, ref_name: str) -> dict:
known = KNOWN_VARIANTS.get(ref_name, {})
annotations = []
if os.path.exists(vcf_path):
with open(vcf_path) as f:
for line in f:
if line.startswith("#"):
continue
parts = line.strip().split("\t")
if len(parts) < 5:
continue
pos = int(parts[1])
ref_base = parts[3]
alt_base = parts[4]
info = {}
for item in parts[7].split(";") if len(parts) > 7 else []:
if "=" in item:
k, v = item.split("=", 1)
info[k] = v
annotation = {
"pos": pos,
"ref": ref_base,
"alt": alt_base,
"depth": int(info.get("DP", "0")),
"freq": float(info.get("AF", "0")),
"gene": "unknown",
"mutation": f"{ref_base}{pos}{alt_base}",
"significance": "Novel variant",
"protein_change": "",
}
if pos in known:
k = known[pos]
annotation["gene"] = k["gene"]
annotation["mutation"] = k["mutation"]
annotation["significance"] = k["significance"]
annotation["protein_change"] = k.get("protein_change", "")
annotations.append(annotation)
return {
"tool": "basic-annotation",
"reference": ref_name,
"annotations": annotations,
"total_annotated": len(annotations),
"known_variants_found": sum(1 for a in annotations if a["significance"] != "Novel variant"),
}
# ---------------------------------------------------------------------------
# Reference genome download
# ---------------------------------------------------------------------------
async def _download_reference(ref_name: str) -> str:
url = REFERENCE_URLS.get(ref_name)
if not url:
raise ValueError(f"Unknown reference genome: {ref_name}")
os.makedirs(CACHE_DIR, exist_ok=True)
fa_path = os.path.join(CACHE_DIR, f"{ref_name}.fa")
if os.path.exists(fa_path) and os.path.getsize(fa_path) > 0:
return fa_path
async with httpx.AsyncClient(timeout=120, follow_redirects=True) as client:
r = await client.get(url)
r.raise_for_status()
data = r.content
if url.endswith(".gz"):
import gzip
data = gzip.decompress(data)
with open(fa_path, "wb") as f:
f.write(data)
return fa_path
# ---------------------------------------------------------------------------
# Synthetic FASTQ generator
# ---------------------------------------------------------------------------
def _generate_synthetic_fastq(ref_seq: str, num_reads: int = 500, read_len: int = 100) -> str:
ref = "".join(line.strip().upper() for line in ref_seq.splitlines() if not line.startswith(">"))
if len(ref) < read_len:
ref = ref * ((read_len // len(ref)) + 1)
lines: list[str] = []
for i in range(num_reads):
start = random.randint(0, len(ref) - read_len)
seq = ref[start:start + read_len]
mut_rate = 0.01
seq = "".join(
random.choice("ACGT") if random.random() < mut_rate else b
for b in seq
)
qual = "".join(chr(33 + min(40, random.randint(20, 40))) for _ in range(read_len))
lines.append(f"@read{i + 1}")
lines.append(seq)
lines.append("+")
lines.append(qual)
return "\n".join(lines)
# ---------------------------------------------------------------------------
# Upload artifacts to Supabase Storage
# ---------------------------------------------------------------------------
def _upload_ngs_files(job_id: str, tmpdir: str, ref_path: str, reference: str) -> dict:
"""Upload BAM, BAI, SAM, VCF, and reference FASTA to Supabase Storage.
Returns dict of URLs keyed by file type."""
from app.services.artifact_storage import _ensure_bucket, BUCKET, get_client
_ensure_bucket()
sb = get_client()
urls = {}
files_to_upload = {
"bam": os.path.join(tmpdir, "sorted.bam"),
"bai": os.path.join(tmpdir, "sorted.bam.bai"),
"sam": os.path.join(tmpdir, "aligned.sam"),
"vcf": os.path.join(tmpdir, "variants.vcf"),
"reference": ref_path,
}
content_types = {
".bam": "application/octet-stream",
".bai": "application/octet-stream",
".sam": "application/octet-stream",
".vcf": "text/vcf",
".fa": "text/plain",
".fasta": "text/plain",
}
for kind, path in files_to_upload.items():
if not os.path.exists(path) or os.path.getsize(path) == 0:
continue
storage_path = f"{job_id}/{kind}"
ext = os.path.splitext(path)[1].lower()
ct = content_types.get(ext, "application/octet-stream")
try:
with open(path, "rb") as f:
data = f.read()
sb.storage.from_(BUCKET).upload(
storage_path, data,
{"content-type": ct, "upsert": "true"},
)
url = sb.storage.from_(BUCKET).get_public_url(storage_path)
urls[kind] = url
logger.info("Uploaded NGS artifact: %s -> %s (%d bytes)", kind, url, len(data))
except Exception as e:
logger.warning("Failed to upload NGS artifact %s: %s", kind, e)
return urls
# ---------------------------------------------------------------------------
# Main Pipeline
# ---------------------------------------------------------------------------
class NGSPipeline(BaseTool):
name = "ngs"
VALID_DEMO = {"synthetic", "demo", "test"}
async def run(self, input: dict) -> dict:
fastq_url = input.get("fastq_url", "").strip()
reference = input.get("reference", "sars-cov-2").strip().lower()
job_id = input.get("job_id", "")
if not fastq_url:
return {"error": "fastq_url is required"}
tmpdir = tempfile.mkdtemp(prefix="ngs_")
steps_completed: list[str] = []
progress: dict[str, str] = {}
try:
# --- Download reference ---
progress["reference"] = "downloading"
ref_path = await asyncio.wait_for(_download_reference(reference), timeout=120)
with open(ref_path) as f:
ref_content = f.read()
# --- Prepare FASTQ ---
fastq_path = os.path.join(tmpdir, "input.fastq")
trimmed_path = os.path.join(tmpdir, "trimmed.fastq")
synthetic = fastq_url.lower() in self.VALID_DEMO
fastq_source = "synthetic"
if synthetic:
fastq_data = _generate_synthetic_fastq(ref_content, num_reads=500, read_len=100)
with open(fastq_path, "w") as f:
f.write(fastq_data)
else:
fastq_source = "url"
try:
await self._download_fastq(fastq_url, fastq_path)
except Exception:
fastq_source = "synthetic"
fastq_data = _generate_synthetic_fastq(ref_content, num_reads=500, read_len=100)
with open(fastq_path, "w") as f:
f.write(fastq_data)
# --- Step 1: Quality Control ---
progress["qc"] = "running"
qc_report_dir = os.path.join(tmpdir, "qc")
os.makedirs(qc_report_dir, exist_ok=True)
qc = await asyncio.to_thread(_run_fastqc, fastq_path, qc_report_dir)
if isinstance(qc, dict) and "error" in qc:
return {"error": qc["error"], "step": "qc", "progress": progress}
steps_completed.append("qc")
progress["qc"] = "done"
# --- Step 2: Trimming ---
progress["trim"] = "running"
trim_report_dir = os.path.join(tmpdir, "trim")
os.makedirs(trim_report_dir, exist_ok=True)
trim_stats = await asyncio.to_thread(_run_fastp, fastq_path, trimmed_path, trim_report_dir)
if not os.path.exists(trimmed_path):
shutil.copy2(fastq_path, trimmed_path)
steps_completed.append("trim")
progress["trim"] = "done"
# --- Step 3: Alignment ---
progress["align"] = "running"
sam_path = os.path.join(tmpdir, "aligned.sam")
try:
align_result = await asyncio.to_thread(_run_alignment, trimmed_path, ref_path, tmpdir)
sam_path = align_result.get("sam_path", sam_path)
except (RuntimeError, FileNotFoundError) as exc:
logger.warning("Native alignment failed (%s) — falling back to Python aligner", exc)
align_result = await asyncio.to_thread(_python_alignment, trimmed_path, ref_path, sam_path)
steps_completed.append("align")
progress["align"] = "done"
# --- Step 4: Variant Calling ---
progress["variants"] = "running"
variant_result = await asyncio.to_thread(_run_variant_calling, sam_path, ref_path, tmpdir)
variants = variant_result.get("variants", [])
vcf_path = variant_result.get("vcf_path", "")
steps_completed.append("variants")
progress["variants"] = "done"
# --- Step 5: Annotation ---
progress["annotate"] = "running"
annotation = await asyncio.to_thread(_run_annotation, vcf_path, reference, tmpdir)
steps_completed.append("annotate")
progress["annotate"] = "done"
# --- Step 6: Upload files for visualization ---
progress["visualization"] = "running"
file_urls = {}
if job_id:
file_urls = await asyncio.to_thread(_upload_ngs_files, job_id, tmpdir, ref_path, reference)
steps_completed.append("visualization")
progress["visualization"] = "done"
# --- Build report ---
report = _build_report(qc, trim_stats, align_result, variants, annotation, reference)
consensus = _build_consensus(ref_content, variants)
return {
"reference": reference,
"reference_size": REFERENCE_SIZES.get(reference, 0),
"fastq_source": fastq_source,
"qc": qc,
"trimming": _summarize_trimming(trim_stats),
"alignment": {
"tool": align_result.get("tool", "unknown"),
"mapped_reads": align_result.get("mapped_reads", 0),
"unmapped_reads": align_result.get("unmapped_reads", 0),
"total_alignments": align_result.get("total_alignments", 0),
"read_region": align_result.get("read_region", ""),
},
"variants": variants[:30],
"annotation": annotation,
"report": report,
"consensus_sequence": f">{reference} consensus (SNVs applied)\n{consensus}",
"file_urls": file_urls,
"steps_completed": steps_completed,
"progress": progress,
"tools_used": {
"qc": qc.get("tool", "unknown"),
"trim": trim_stats.get("tool", "unknown"),
"align": align_result.get("tool", "unknown"),
"variant": variant_result.get("tool", "unknown"),
"annotate": annotation.get("tool", "unknown"),
},
}
except ValueError as e:
return {"error": str(e), "progress": progress}
except httpx.HTTPStatusError as e:
return {"error": f"Download failed (HTTP {e.response.status_code})", "progress": progress}
except asyncio.TimeoutError:
return {"error": "Pipeline timed out", "progress": progress}
except Exception as e:
logger.exception("NGS pipeline failed")
return {"error": f"Pipeline failed: {e}", "progress": progress}
finally:
shutil.rmtree(tmpdir, ignore_errors=True)
async def _download_fastq(self, url: str, dest: str) -> str:
async with httpx.AsyncClient(timeout=120, follow_redirects=True) as client:
async with client.stream("GET", url) as r:
r.raise_for_status()
with open(dest, "wb") as f:
async for chunk in r.aiter_bytes():
f.write(chunk)
return dest
def _summarize_trimming(trim_stats: dict) -> dict:
before = trim_stats.get("before_filtering", {})
after = trim_stats.get("after_filtering", {})
return {
"tool": trim_stats.get("tool", "unknown"),
"reads_before": before.get("total_reads", 0),
"reads_after": after.get("total_reads", 0),
"reads_discarded": trim_stats.get("filtering_result", {}).get("low_quality_reads", 0),
}
def _build_consensus(reference_seq: str, variants: list[dict]) -> str:
ref_lines = reference_seq.splitlines()
ref = "".join(line.strip().upper() for line in ref_lines if not line.startswith(">"))
seq = list(ref)
for v in variants:
pos = v.get("pos", 0) - 1
alt = v.get("alt", "")
if 0 <= pos < len(seq):
seq[pos] = alt
return "".join(seq)
def _build_report(qc: dict, trim: dict, align: dict, variants: list[dict], annotation: dict, ref_name: str) -> dict:
total_variants = len(variants)
snv_count = sum(1 for v in variants if len(v.get("ref", "")) == 1 and len(v.get("alt", "")) == 1)
known_count = annotation.get("known_variants_found", 0)
novel_count = total_variants - known_count
return {
"reference": ref_name,
"steps": ["QC", "Trimming", "Alignment", "Variant Calling", "Annotation", "Visualization"],
"qc_summary": {
"total_reads": qc.get("total_reads", 0),
"total_bases": qc.get("total_bases", 0),
"mean_quality": qc.get("mean_quality", 0),
"q30_percent": qc.get("q30_percent", 0),
"gc_percent": qc.get("gc_percent", 0),
},
"trimming_summary": {
"reads_before": trim.get("reads_before", 0),
"reads_after": trim.get("reads_after", 0),
},
"alignment_summary": {
"mapped_reads": align.get("mapped_reads", 0),
"unmapped_reads": align.get("unmapped_reads", 0),
"mapping_rate": round(
align.get("mapped_reads", 0) / max(align.get("total_alignments", 1), 1) * 100, 1
),
},
"variant_summary": {
"total_variants": total_variants,
"snv_count": snv_count,
"known_variants": known_count,
"novel_variants": novel_count,
},
}