use anyhow::{Context, Result};
use clap::{ArgGroup, Parser};
use lexical_core::parse;
use memchr::memchr;
use memmap2::Mmap;
use rayon::prelude::*;
use rustc_hash::{FxHashMap, FxHashSet};
use std::{
fs::File,
io::{self, BufWriter, IoSlice, Write},
path::{Path, PathBuf},
};
use crate::{
CommonArgs, Interval, TreeIndexData, load_gof, write_gff_output,
};
const MISSING: u64 = u64::MAX;
const IOV_BATCH: usize = 256;
const WRITE_BUF_SIZE: usize = 32 * 1024 * 1024;
#[derive(Debug, Clone)]
pub struct RootMatched {
pub root: u32,
pub matched: Vec<u32>,
}
#[derive(Parser, Debug)]
#[command(
about = "Extract models by a region or regions from a BED file",
long_about = "This tool extracts features and their parent models that intersect with specified regions"
)]
#[clap(group(
ArgGroup::new("regions").required(true).args(&["region", "bed"])
))]
#[clap(group(
ArgGroup::new("mode").args(&["contained", "contains_region", "overlap"])
))]
pub struct IntersectArgs {
#[clap(flatten)]
pub common: CommonArgs,
#[arg(short = 'r', long, group = "regions")]
pub region: Option<String>,
#[arg(short = 'b', long, group = "regions")]
pub bed: Option<PathBuf>,
#[arg(short = 'c', long, group = "mode")]
pub contained: bool,
#[arg(short = 'C', long, group = "mode")]
pub contains_region: bool,
#[arg(short = 'O', long, group = "mode")]
pub overlap: bool,
#[arg(short = 'I', long, default_value_t = false)]
pub invert: bool,
}
#[derive(Debug, Clone, Copy)]
pub enum OverlapMode {
Contained,
ContainsRegion,
Overlap,
}
pub fn gff_type_allowed(line: &[u8], allow: &FxHashSet<String>) -> bool {
let mut off = 0usize;
let mut tabs = 0u8;
while tabs < 2 {
match memchr(b'\t', &line[off..]) {
Some(i) => {
off += i + 1;
tabs += 1;
}
None => return false,
}
}
let i2 = match memchr(b'\t', &line[off..]) {
Some(i) => off + i,
None => return false,
};
let ty = &line[off..i2];
match std::str::from_utf8(ty) {
Ok(s) => allow.contains(s),
Err(_) => false,
}
}
pub fn query_features(
index_data: &TreeIndexData,
regions: &[(u32, u32, u32)],
mode: OverlapMode,
invert: bool,
verbose: bool,
) -> Result<Vec<(u32, u32, u32)>> {
let buckets: Vec<Vec<(u32, u32, u32)>> = {
let mut b = vec![Vec::new(); index_data.seqid_to_num.len()];
for (chr, start, end) in regions.iter().copied() {
b[chr as usize].push((chr, start, end));
}
b
};
let mut results = Vec::new();
{
for (&seq_num, tree) in &index_data.chr_entries {
let chr_regs = &buckets[seq_num as usize];
if chr_regs.is_empty() {
continue;
}
if verbose {
eprintln!(
"[DEBUG] Querying chromosome {} with {} regions",
seq_num,
chr_regs.len()
);
}
let mut hits: Vec<&Interval<u32>> = Vec::new();
for &(_, rstart, rend) in chr_regs {
hits.clear();
tree.query_interval(rstart, rend, &mut hits);
for &iv in &hits {
let keep = match mode {
OverlapMode::Contained => {
iv.start >= rstart && iv.end <= rend
}
OverlapMode::ContainsRegion => {
iv.start <= rstart && iv.end >= rend
}
OverlapMode::Overlap => {
true
}
};
if invert ^ keep {
results.push((iv.root_fid, iv.start, iv.end));
}
}
}
}
}
Ok(results)
}
pub fn parse_region(
region: &str,
seqid_map: &FxHashMap<String, u32>,
common: &CommonArgs,
) -> Result<(u32, u32, u32)> {
let (seq, range) = region
.split_once(':')
.context("Invalid region format, expected 'chr:start-end'")?;
let (s, e) = range
.split_once('-')
.context("Invalid range format, expected 'start-end'")?;
let start = s.parse::<u32>()?;
let end = e.parse::<u32>()?;
let chr = seqid_map
.get(seq)
.with_context(|| format!("Sequence ID not found: {}", seq))?;
if start >= end {
anyhow::bail!("Region start must be less than end ({} >= {})", start, end);
}
if common.verbose {
eprintln!(
"[DEBUG] Parsed region: chr={}, start={}, end={}",
chr, start, end
);
}
Ok((*chr, start, end))
}
pub fn parse_bed_file(
bed_path: &Path,
seqid_map: &FxHashMap<String, u32>,
) -> Result<Vec<(u32, u32, u32)>> {
let mmap = {
let file = File::open(bed_path)?;
unsafe { Mmap::map(&file)? }
};
let regions = {
let mut regions = Vec::new();
for line in mmap.split(|&b| b == b'\n') {
if line.is_empty() || line[0] == b'#' {
continue;
}
let line_str = std::str::from_utf8(line)?;
let mut parts = line_str.split_ascii_whitespace();
let (Some(seq), Some(s), Some(e)) = (parts.next(), parts.next(), parts.next()) else {
continue;
};
let Some(&chr) = seqid_map.get(seq) else {
continue;
};
let start = parse::<u32>(s.as_bytes())?;
let end = parse::<u32>(e.as_bytes())?;
regions.push((chr, start, end));
}
regions
};
Ok(regions)
}
pub fn write_gff_match_only_by_coords(
gff_path: &Path,
blocks: &[(u32, u64, u64)], query_ivmap: &FxHashMap<String, Vec<(u32, u32)>>,
types_filter: Option<&str>,
output_path: &Option<PathBuf>,
mode: OverlapMode,
verbose: bool,
) -> Result<()> {
let (mmap, file_len) = {
let file = std::fs::File::open(gff_path)
.with_context(|| format!("Cannot open GFF: {:?}", gff_path))?;
let mmap = unsafe { Mmap::map(&file) }
.with_context(|| format!("mmap failed for {:?}", gff_path))?;
let len = mmap.len();
(mmap, len)
};
let type_allow: Option<FxHashSet<String>> = {
types_filter.map(|s| {
s.split(',')
.map(|t| t.trim().to_string())
.filter(|t| !t.is_empty())
.collect()
})
};
let mut parts: Vec<(u64, Vec<(u64, u64)>)> = {
let bytes_out = std::sync::atomic::AtomicU64::new(0);
let parts: Vec<(u64, Vec<(u64, u64)>)> = blocks
.par_iter()
.filter_map(|&(root, start, end)| {
if start == MISSING {
eprintln!("[WARN] skipped fid={} due to sentinel start offset", root);
return None;
}
let s = start as usize;
let e = (end as usize).min(file_len);
if s >= e || e > file_len {
return None;
}
let src = &mmap[s..e];
let mut matched_offsets: Vec<(u64, u64)> = Vec::with_capacity(256);
let mut pos = 0usize;
while pos < src.len() {
let nl = match memchr(b'\n', &src[pos..]) {
Some(i) => pos + i + 1, None => src.len(),
};
let line = &src[pos..nl];
let line_nocr = if line.ends_with(b"\n") {
&line[..line.len() - 1]
} else {
line
};
if !line_nocr.is_empty() && line_nocr[0] != b'#' {
let mut pass = true;
if let Some(allow) = &type_allow
&& !gff_type_allowed(line_nocr, allow)
{
pass = false;
}
if pass && gff_line_overlaps_queries(line_nocr, query_ivmap, mode) {
let abs_start = start + pos as u64;
let abs_end = start + nl as u64;
matched_offsets.push((abs_start, abs_end));
bytes_out.fetch_add(
(abs_end - abs_start) as u64,
std::sync::atomic::Ordering::Relaxed,
);
}
}
pos = nl;
}
if matched_offsets.is_empty() {
None
} else {
Some((start, matched_offsets))
}
})
.collect();
parts
};
{
parts.sort_unstable_by_key(|(s, _)| *s);
}
fn write_all_vectored<W: Write>(w: &mut W, mut slices: Vec<&[u8]>) -> io::Result<()> {
if slices.is_empty() {
return Ok(());
}
while !slices.is_empty() {
let iov: Vec<IoSlice<'_>> = slices.iter().map(|s| IoSlice::new(s)).collect();
let wrote = w.write_vectored(&iov)?;
if wrote == 0 {
return Err(io::Error::new(
io::ErrorKind::WriteZero,
"write_vectored returned 0",
));
}
let mut remaining = wrote;
let mut drop_count = 0;
for s in &mut slices {
if remaining == 0 {
break;
}
if remaining >= s.len() {
remaining -= s.len();
drop_count += 1;
} else {
*s = &s[remaining..];
remaining = 0;
}
}
if drop_count > 0 {
slices.drain(0..drop_count);
}
}
Ok(())
}
{
let mut batch: Vec<&[u8]> = Vec::with_capacity(IOV_BATCH);
if let Some(p) = output_path {
let file = std::fs::File::create(p)?;
let mut writer = BufWriter::with_capacity(WRITE_BUF_SIZE, file);
for (_, ranges) in parts.iter() {
for &(ls, le) in ranges {
let slice = &mmap[ls as usize..le as usize];
batch.push(slice);
if batch.len() >= IOV_BATCH {
write_all_vectored(&mut writer, std::mem::take(&mut batch))?;
}
}
}
if !batch.is_empty() {
write_all_vectored(&mut writer, std::mem::take(&mut batch))?;
}
writer.flush()?;
} else {
let stdout = std::io::stdout();
let handle = stdout.lock();
let mut writer = BufWriter::with_capacity(WRITE_BUF_SIZE, handle);
for (_, ranges) in parts.iter() {
for &(ls, le) in ranges {
let slice = &mmap[ls as usize..le as usize];
batch.push(slice);
if batch.len() >= IOV_BATCH {
write_all_vectored(&mut writer, std::mem::take(&mut batch))?;
}
}
}
if !batch.is_empty() {
write_all_vectored(&mut writer, std::mem::take(&mut batch))?;
}
writer.flush()?;
}
}
if verbose {
eprintln!(
"[INFO] match-only by coords completed; minput blocks {}",
blocks.len()
);
}
Ok(())
}
pub fn gff_line_overlaps_queries(
line: &[u8],
ivmap: &FxHashMap<String, Vec<(u32, u32)>>,
mode: OverlapMode,
) -> bool {
let mut off = 0usize;
let i1 = match memchr(b'\t', &line[off..]) {
Some(i) => off + i,
None => return false,
};
let seq = &line[off..i1];
off = i1 + 1;
let i2 = match memchr(b'\t', &line[off..]) {
Some(i) => off + i,
None => return false,
};
off = i2 + 1;
let i3 = match memchr(b'\t', &line[off..]) {
Some(i) => off + i,
None => return false,
};
off = i3 + 1;
let i4 = match memchr(b'\t', &line[off..]) {
Some(i) => off + i,
None => return false,
};
let start = match parse_u32_ascii(&line[off..i4]) {
Some(v) => v,
None => return false,
};
off = i4 + 1;
let i5 = match memchr(b'\t', &line[off..]) {
Some(i) => off + i,
None => return false,
};
let end = match parse_u32_ascii(&line[off..i5]) {
Some(v) => v,
None => return false,
};
let seq_str = match std::str::from_utf8(seq) {
Ok(s) => s,
Err(_) => return false,
};
let ivs = match ivmap.get(seq_str) {
Some(v) => v,
None => return false,
};
for &(qs, qe) in ivs {
let keep = match mode {
OverlapMode::Contained => {
start >= qs && end <= qe
}
OverlapMode::ContainsRegion => {
start <= qs && end >= qe
}
OverlapMode::Overlap => {
(qs <= start && start <= qe)
|| (qs <= end && end <= qe)
|| (start <= qs && qs <= end)
|| (start <= qe && qe <= end)
}
};
if keep {
return true;
}
}
false
}
#[inline]
fn parse_u32_ascii(s: &[u8]) -> Option<u32> {
let mut v: u32 = 0;
if s.is_empty() {
return None;
}
for &c in s {
if !c.is_ascii_digit() {
return None;
}
v = v.checked_mul(10)?.checked_add((c - b'0') as u32)?;
}
Some(v)
}
pub fn run(args: &IntersectArgs) -> Result<()> {
let verbose = args.common.verbose;
if verbose {
eprintln!("[DEBUG] Starting processing of {:?}", args.common.input);
eprintln!(
"[DEBUG] Thread pool initialized with {} threads",
args.common.effective_threads()
);
}
let mode = if args.contained {
OverlapMode::Contained
} else if args.contains_region {
OverlapMode::ContainsRegion
} else {
OverlapMode::Overlap
};
let index_data = TreeIndexData::load_tree_index(&args.common.input)?;
let seqid_map = &index_data.seqid_to_num;
let regions = {
if let Some(bed) = &args.bed {
parse_bed_file(bed, seqid_map)?
} else if let Some(r) = &args.region {
vec![parse_region(r, seqid_map, &args.common)?]
} else {
anyhow::bail!("No region specified");
}
};
if verbose {
eprintln!(
"[DEBUG] Starting query_features with {} regions",
regions.len()
);
eprintln!(
"[DEBUG] Mode: {:?}",
mode
);
}
let feats = {
query_features(
&index_data,
®ions,
mode,
args.invert,
args.common.verbose,
)?
};
let gof = load_gof(&args.common.input)?;
let root_matches: Vec<RootMatched> = {
let mut grouped: FxHashMap<u32, Vec<u32>> = FxHashMap::default();
for (root, _s, _e) in feats {
grouped.entry(root).or_default().push(root);
}
grouped
.into_iter()
.map(|(root, matched)| RootMatched { root, matched })
.collect()
};
let roots: Vec<u32> = {
let mut s: FxHashSet<u32> = FxHashSet::default();
for rm in &root_matches {
s.insert(rm.root);
}
s.into_iter().collect()
};
let blocks: Vec<(u32, u64, u64)> = gof.roots_to_offsets(&roots, args.common.effective_threads());
if !args.common.entire_group || args.common.types.is_some() {
let query_ivmap: FxHashMap<String, Vec<(u32, u32)>> = {
let mut num_to_seq: FxHashMap<u32, String> = FxHashMap::default();
for (name, &num) in index_data.seqid_to_num.iter() {
num_to_seq.insert(num, name.clone());
}
let mut m: FxHashMap<String, Vec<(u32, u32)>> = FxHashMap::default();
for &(chr_num, s, e) in ®ions {
if let Some(seq_name) = num_to_seq.get(&chr_num) {
m.entry(seq_name.clone()).or_default().push((s, e));
}
}
m
};
{
write_gff_match_only_by_coords(
args.common.input.as_path(),
&blocks,
&query_ivmap,
args.common.types.as_deref(),
&args.common.output,
mode,
args.common.verbose,
)?;
}
} else {
write_gff_output(
args.common.input.as_path(),
&blocks,
&args.common.output,
args.common.verbose,
)?;
}
Ok(())
}