use std::cell::Cell;
use std::io;
use std::iter::repeat_n;
use std::mem::MaybeUninit;
use std::slice;
use crate::TwobitSequenceData;
use memmap2::Mmap;
use pyo3::exceptions::{PyKeyError, PyValueError};
use pyo3::ffi::{Py_DECREF, PyObject, PyUnicode_1BYTE_DATA, PyUnicode_CheckExact, PyUnicode_New};
use pyo3::prelude::*;
use pyo3::pybacked::PyBackedStr;
use pyo3::types::{PyList, PyNone, PyString};
use scopeguard::ScopeGuard;
#[pyclass(frozen)]
#[derive(Debug)]
pub struct TwobitReader {
inner: crate::TwobitReader,
}
#[pymethods]
impl TwobitReader {
#[staticmethod]
pub fn open(path: &str) -> io::Result<Self> {
crate::TwobitReader::open(path).map(|inner| TwobitReader { inner })
}
#[staticmethod]
pub fn open_masked(path: &str) -> io::Result<Self> {
crate::TwobitReader::open_masked(path).map(|inner| TwobitReader { inner })
}
pub fn names(&self) -> Vec<String> {
self.inner.iter_names().map(|s| s.to_string()).collect()
}
pub fn contains_name(&self, name: &str) -> bool {
self.inner.contains_name(name)
}
pub fn seq_len(&self, name: &str) -> PyResult<usize> {
Ok(self.seq_data(name)?.dna_len)
}
pub fn get<'py>(&self, py: Python<'py>, name: &str, start: usize, end: usize) -> PyResult<Bound<'py, PyString>> {
let seq = self.seq_data(name)?;
check_range(seq, start, end)?;
let (unicode_guard, buf) = new_unicode_guarded(py, end - start)?;
decode_and_fill_blocks(&self.inner.mmap, seq, start, buf);
Ok(as_bound_unicode(py, unicode_guard))
}
pub fn get_batch<'py>(&self, py: Python<'py>, batch: &Bound<'_, PyAny>) -> PyResult<Bound<'py, PyList>> {
self.get_batch_impl(py, batch, 0)
}
#[rustfmt::skip]
pub fn get_inclusive<'py>(&self, py: Python<'py>, name: &str, start: usize, end: usize) -> PyResult<Bound<'py, PyString>> {
if start < 1 {
return Err(PyValueError::new_err("Invalid start for inclusive range (start == 0)"));
}
self.get(py, name, start - 1, end)
}
#[rustfmt::skip]
pub fn get_batch_inclusive<'py>(&self, py: Python<'py>, batch: &Bound<'_, PyAny>) -> PyResult<Bound<'py, PyList>> {
self.get_batch_impl(py, batch, 1)
}
#[rustfmt::skip]
pub fn concat<'py>(&self, py: Python<'py>, name: &str, ranges: &Bound<'_, PyAny>) -> PyResult<Bound<'py, PyString>> {
self.concat_impl(py, name, ranges, 0)
}
pub fn concat_inclusive<'py>(
&self,
py: Python<'py>,
name: &str,
ranges: &Bound<'_, PyAny>,
) -> PyResult<Bound<'py, PyString>> {
self.concat_impl(py, name, ranges, 1)
}
pub fn prefetch(&self, batch: &Bound<'_, PyAny>) -> PyResult<()> {
self.prefetch_impl(batch, 0)
}
pub fn prefetch_inclusive(&self, batch: &Bound<'_, PyAny>) -> PyResult<()> {
self.prefetch_impl(batch, 1)
}
}
impl TwobitReader {
#[inline]
fn seq_data(&self, name: &str) -> PyResult<&TwobitSequenceData> {
self.inner.try_get_seq_data_by_name(name).ok_or_else(|| PyKeyError::new_err(name.to_string()))
}
fn prefetch_impl(&self, batch: &Bound<'_, PyAny>, base: usize) -> PyResult<()> {
let ec = PyErrCapture::default();
let ranges = batch.try_iter()?.map_while(|item| {
let (name, start, end) = ec.ok(ec.ok(item)?.extract::<(PyBackedStr, usize, usize)>())?;
ec.ok(self.check_prefetch_range(&name, start, end, base))?;
Some((name, start, end))
});
if base == 1 {
self.inner.prefetch_inclusive(ranges);
} else {
self.inner.prefetch(ranges);
}
ec.into_result()
}
#[inline]
fn check_prefetch_range(&self, name: &str, start: usize, end: usize, base: usize) -> PyResult<()> {
self.seq_data(name)?;
check_start_inclusive(start, base)?;
if start - base > end {
return Err(PyValueError::new_err("invalid range (start > end)"));
}
Ok(())
}
#[rustfmt::skip]
fn get_batch_impl<'py>(&self, py: Python<'py>, batch: &Bound<'_, PyAny>, base: usize) -> PyResult<Bound<'py, PyList>> {
let get = |item: Result<Bound<'_, PyAny>, PyErr>| -> Result<Bound<'py, PyString>, PyErr> {
let item = item?;
let (name, start, end): (&str, usize, usize) = item.extract()?;
if base == 1 {
self.get_inclusive(py, name, start, end)
} else {
self.get(py, name, start, end)
}
};
if let Ok(len) = batch.len() {
let list = PyList::new(py, repeat_n(PyNone::get(py), len))?;
for (i, item) in batch.try_iter()?.enumerate() {
list.set_item(i, get(item)?)?;
}
Ok(list)
} else {
let list = PyList::empty(py);
for item in batch.try_iter()? {
list.append(get(item)?)?;
}
Ok(list)
}
}
#[rustfmt::skip]
fn concat_impl<'py>(&self, py: Python<'py>, name: &str, ranges: &Bound<'_, PyAny>, base: usize) -> PyResult<Bound<'py, PyString>> {
let mut range_iter = ranges.try_iter()?;
if ranges.is(&range_iter) {
let ec = PyErrCapture::default();
let result = self.inner.concat_iter(
name,
ranges.try_iter()?.map_while(|item| {
let (start, end) = ec.ok(ec.ok(item)?.extract::<(usize, usize)>())?;
ec.ok(check_start_inclusive(start, base))?;
Some((start - base, end))
}),
);
ec.check()?;
let (unicode_guard, buf) = new_unicode_guarded(py, result.len())?;
buf.write_copy_of_slice(result.as_bytes());
Ok(as_bound_unicode(py, unicode_guard))
} else {
let seq = self.seq_data(name)?;
let total_len = range_iter.try_fold(0, |total, item| -> PyResult<usize> {
let (start, end) = item?.extract()?;
check_start_inclusive(start, base)?; check_range(seq, start - base, end)?;
Ok(total + end - (start - base))
})?;
let (unicode_guard, mut buf) = new_unicode_guarded(py, total_len)?;
for item in ranges.try_iter()? {
let (start, end): (usize, usize) = item?.extract()?;
check_start_inclusive(start, base)?; check_range(seq, start - base, end)?;
let start = start - base;
if end - start > buf.len() {
return Err(PyValueError::new_err("ranges grew between iterations"));
}
let (head, tail) = buf.split_at_mut(end - start);
decode_and_fill_blocks(&self.inner.mmap, seq, start, head);
buf = tail;
}
if !buf.is_empty() {
return Err(PyValueError::new_err("ranges shrank between iterations"));
}
Ok(as_bound_unicode(py, unicode_guard))
}
}
}
#[pymodule(gil_used = false)]
fn twobitreader_rs(m: &Bound<'_, PyModule>) -> PyResult<()> {
m.add_class::<TwobitReader>()?;
m.add_function(wrap_pyfunction!(reverse_complement, m)?)?;
Ok(())
}
#[pyfunction]
fn reverse_complement(dna: String) -> PyResult<String> {
Ok(crate::reverse_complement(dna))
}
#[inline]
#[allow(clippy::type_complexity)]
fn new_unicode_guarded<'py>(
py: Python<'py>,
len: usize,
) -> PyResult<(ScopeGuard<*mut PyObject, impl FnOnce(*mut PyObject)>, &'py mut [MaybeUninit<u8>])> {
if len >= isize::MAX as usize {
return Err(PyValueError::new_err("range length exceeded string size limit"));
}
let unicode = unsafe { PyUnicode_New(len as isize, 127) };
if unicode.is_null() {
return Err(PyErr::fetch(py));
}
let unicode_guard = scopeguard::guard(unicode, |o| {
unsafe { Py_DECREF(o) };
});
let buf = unsafe { slice::from_raw_parts_mut(PyUnicode_1BYTE_DATA(unicode).cast(), len) };
Ok((unicode_guard, buf))
}
#[inline]
#[allow(clippy::type_complexity)]
fn as_bound_unicode<'py>(
py: Python<'py>,
unicode_guard: ScopeGuard<*mut PyObject, impl FnOnce(*mut PyObject)>,
) -> Bound<'py, PyString> {
debug_assert!(!unicode_guard.is_null());
debug_assert!(unsafe { PyUnicode_CheckExact(*unicode_guard) != 0 });
let unicode = ScopeGuard::into_inner(unicode_guard);
unsafe { Bound::from_owned_ptr(py, unicode).cast_into_unchecked() }
}
#[inline]
fn check_range(seq: &crate::TwobitSequenceData, start: usize, end: usize) -> PyResult<()> {
if start > end {
return Err(PyValueError::new_err("invalid range (start > end)"));
}
if end > seq.dna_len {
return Err(PyValueError::new_err("invalid end (end > dna_len)"));
}
Ok(())
}
#[inline]
fn check_start_inclusive(start: usize, base: usize) -> PyResult<()> {
if start < base {
debug_assert_eq!(base, 1, "expected base=1 but found {base}");
return Err(PyValueError::new_err("invalid start (0) for 1-based range"));
}
Ok(())
}
#[inline]
fn decode_and_fill_blocks(mmap: &Mmap, seq: &TwobitSequenceData, start: usize, buf: &mut [MaybeUninit<u8>]) {
crate::decode_from_mmap(mmap, seq, start, buf);
crate::fill_blocks(seq, start, unsafe { buf.assume_init_mut() });
}
#[derive(Default)]
struct PyErrCapture(Cell<Option<PyErr>>);
impl PyErrCapture {
fn ok<T>(&self, result: PyResult<T>) -> Option<T> {
match result {
Ok(value) => Some(value),
Err(err) => {
let first = self.0.take().or(Some(err));
self.0.set(first);
None
}
}
}
fn check(&self) -> PyResult<()> {
self.0.take().map_or(Ok(()), Err)
}
fn into_result(self) -> PyResult<()> {
self.0.into_inner().map_or(Ok(()), Err)
}
}