use std::future::Future;
use std::ops::Range;
use std::sync::Arc;
use bytes::Bytes;
use futures::future::BoxFuture;
use futures::{FutureExt, TryFutureExt};
use tokio::runtime::Handle;
use crate::arrow::arrow_reader::ArrowReaderOptions;
use crate::arrow::async_reader::{AsyncFileReader, MetadataSuffixFetch};
use crate::errors::{ParquetError, Result};
use crate::file::metadata::ParquetMetaData;
#[derive(Clone, Debug)]
pub struct SpawnedReader<R> {
inner: R,
handle: Handle,
}
impl<R> SpawnedReader<R> {
pub fn new(inner: R, handle: Handle) -> Self {
Self { inner, handle }
}
pub fn into_inner(self) -> R {
self.inner
}
}
fn spawn<T>(
handle: &Handle,
fut: impl Future<Output = Result<T>> + Send + 'static,
) -> BoxFuture<'static, Result<T>>
where
T: Send + 'static,
{
handle
.spawn(fut)
.map_ok_or_else(
|e| match e.try_into_panic() {
Err(e) => Err(ParquetError::External(Box::new(e))),
Ok(p) => std::panic::resume_unwind(p),
},
|res| res,
)
.boxed()
}
impl<R> AsyncFileReader for SpawnedReader<R>
where
R: AsyncFileReader + Clone + Send + 'static,
{
fn get_bytes(&mut self, range: Range<u64>) -> BoxFuture<'_, Result<Bytes>> {
let mut inner = self.inner.clone();
spawn(&self.handle, async move { inner.get_bytes(range).await })
}
fn get_byte_ranges(&mut self, ranges: Vec<Range<u64>>) -> BoxFuture<'_, Result<Vec<Bytes>>> {
let mut inner = self.inner.clone();
spawn(
&self.handle,
async move { inner.get_byte_ranges(ranges).await },
)
}
fn get_metadata<'a>(
&'a mut self,
options: Option<&'a ArrowReaderOptions>,
) -> BoxFuture<'a, Result<Arc<ParquetMetaData>>> {
let mut inner = self.inner.clone();
let options = options.cloned();
spawn(&self.handle, async move {
inner.get_metadata(options.as_ref()).await
})
}
}
impl<R> MetadataSuffixFetch for &mut SpawnedReader<R>
where
R: AsyncFileReader + Clone + Send + 'static,
for<'a> &'a mut R: MetadataSuffixFetch,
{
fn fetch_suffix(&mut self, suffix: usize) -> BoxFuture<'_, Result<Bytes>> {
let mut inner = self.inner.clone();
spawn(&self.handle, async move {
(&mut inner).fetch_suffix(suffix).await
})
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::arrow::ParquetRecordBatchStreamBuilder;
use crate::file::metadata::ParquetMetaDataReader;
use futures::TryStreamExt;
use std::thread::ThreadId;
#[derive(Clone)]
struct InMemoryReader {
data: Bytes,
threads: Arc<std::sync::Mutex<Vec<ThreadId>>>,
}
impl InMemoryReader {
fn new(data: Bytes) -> Self {
Self {
data,
threads: Default::default(),
}
}
}
impl AsyncFileReader for InMemoryReader {
fn get_bytes(&mut self, range: Range<u64>) -> BoxFuture<'_, Result<Bytes>> {
self.threads
.lock()
.unwrap()
.push(std::thread::current().id());
let data = self.data.slice(range.start as usize..range.end as usize);
futures::future::ready(Ok(data)).boxed()
}
fn get_metadata<'a>(
&'a mut self,
options: Option<&'a ArrowReaderOptions>,
) -> BoxFuture<'a, Result<Arc<ParquetMetaData>>> {
self.threads
.lock()
.unwrap()
.push(std::thread::current().id());
let metadata = ParquetMetaDataReader::new()
.with_arrow_reader_options(options)
.parse_and_finish(&self.data);
futures::future::ready(metadata.map(Arc::new)).boxed()
}
}
#[tokio::test]
async fn test_spawned_reader() {
let testdata = arrow::util::test_util::parquet_test_data();
let path = format!("{testdata}/alltypes_plain.parquet");
let data = Bytes::from(std::fs::read(path).unwrap());
let rt = tokio::runtime::Builder::new_multi_thread()
.worker_threads(1)
.build()
.unwrap();
let inner = InMemoryReader::new(data);
let threads = inner.threads.clone();
let reader = SpawnedReader::new(inner, rt.handle().clone());
let builder = ParquetRecordBatchStreamBuilder::new(reader).await.unwrap();
let batches: Vec<_> = builder.build().unwrap().try_collect().await.unwrap();
assert_eq!(batches.len(), 1);
assert_eq!(batches[0].num_rows(), 8);
let current_id = std::thread::current().id();
let threads = threads.lock().unwrap();
assert!(!threads.is_empty());
assert!(threads.iter().all(|id| *id != current_id));
tokio::runtime::Handle::current().spawn_blocking(move || drop(rt));
}
#[tokio::test]
async fn test_spawned_reader_fails_on_shutdown_runtime() {
let rt = tokio::runtime::Builder::new_multi_thread()
.worker_threads(1)
.build()
.unwrap();
let inner = InMemoryReader::new(Bytes::from_static(b"PAR1"));
let mut reader = SpawnedReader::new(inner, rt.handle().clone());
rt.shutdown_background();
let err = reader.get_bytes(0..1).await.unwrap_err().to_string();
assert!(err.contains("was cancelled"), "{err}");
}
}