arrow_avro/reader/async_reader/
spawn.rs1use std::future::Future;
19use std::ops::Range;
20
21use bytes::Bytes;
22use futures::future::BoxFuture;
23use futures::{FutureExt, TryFutureExt};
24use tokio::runtime::Handle;
25
26use crate::errors::AvroError;
27use crate::reader::async_reader::AsyncFileReader;
28
29#[derive(Clone, Debug)]
45pub struct SpawnedReader<R> {
46 inner: R,
47 handle: Handle,
48}
49
50impl<R> SpawnedReader<R> {
51 pub fn new(inner: R, handle: Handle) -> Self {
53 Self { inner, handle }
54 }
55
56 pub fn into_inner(self) -> R {
58 self.inner
59 }
60}
61
62fn spawn<T>(
65 handle: &Handle,
66 fut: impl Future<Output = Result<T, AvroError>> + Send + 'static,
67) -> BoxFuture<'static, Result<T, AvroError>>
68where
69 T: Send + 'static,
70{
71 handle
72 .spawn(fut)
73 .map_ok_or_else(
74 |e| match e.try_into_panic() {
75 Err(e) => Err(AvroError::External(Box::new(e))),
76 Ok(p) => std::panic::resume_unwind(p),
77 },
78 |res| res,
79 )
80 .boxed()
81}
82
83impl<R> AsyncFileReader for SpawnedReader<R>
84where
85 R: AsyncFileReader + Clone + Send + 'static,
86{
87 fn get_bytes(&mut self, range: Range<u64>) -> BoxFuture<'_, Result<Bytes, AvroError>> {
88 let mut inner = self.inner.clone();
89 spawn(&self.handle, async move { inner.get_bytes(range).await })
90 }
91
92 fn get_byte_ranges(
93 &mut self,
94 ranges: Vec<Range<u64>>,
95 ) -> BoxFuture<'_, Result<Vec<Bytes>, AvroError>> {
96 let mut inner = self.inner.clone();
97 spawn(
98 &self.handle,
99 async move { inner.get_byte_ranges(ranges).await },
100 )
101 }
102}
103
104#[cfg(test)]
105mod tests {
106 use super::*;
107 use std::sync::{Arc, Mutex};
108 use std::thread::ThreadId;
109
110 #[derive(Clone)]
112 struct InMemoryReader {
113 data: Bytes,
114 threads: Arc<Mutex<Vec<ThreadId>>>,
115 }
116
117 impl AsyncFileReader for InMemoryReader {
118 fn get_bytes(&mut self, range: Range<u64>) -> BoxFuture<'_, Result<Bytes, AvroError>> {
119 self.threads
120 .lock()
121 .unwrap()
122 .push(std::thread::current().id());
123 let data = self.data.slice(range.start as usize..range.end as usize);
124 futures::future::ready(Ok(data)).boxed()
125 }
126 }
127
128 #[tokio::test]
129 async fn test_spawned_reader() {
130 let rt = tokio::runtime::Builder::new_multi_thread()
131 .worker_threads(1)
132 .build()
133 .unwrap();
134
135 let inner = InMemoryReader {
136 data: Bytes::from_static(b"hello world"),
137 threads: Default::default(),
138 };
139 let threads = inner.threads.clone();
140 let mut reader = SpawnedReader::new(inner, rt.handle().clone());
141
142 let bytes = reader.get_bytes(0..5).await.unwrap();
143 assert_eq!(bytes.as_ref(), b"hello");
144
145 let ranges = reader.get_byte_ranges(vec![0..5, 6..11]).await.unwrap();
146 assert_eq!(ranges[1].as_ref(), b"world");
147
148 let current_id = std::thread::current().id();
150 let threads = threads.lock().unwrap();
151 assert!(!threads.is_empty());
152 assert!(threads.iter().all(|id| *id != current_id));
153
154 tokio::runtime::Handle::current().spawn_blocking(move || drop(rt));
156 }
157
158 #[tokio::test]
159 async fn test_spawned_reader_fails_on_shutdown_runtime() {
160 let rt = tokio::runtime::Builder::new_multi_thread()
161 .worker_threads(1)
162 .build()
163 .unwrap();
164
165 let inner = InMemoryReader {
166 data: Bytes::from_static(b"hello world"),
167 threads: Default::default(),
168 };
169 let mut reader = SpawnedReader::new(inner, rt.handle().clone());
170
171 rt.shutdown_background();
172
173 let err = reader.get_bytes(0..1).await.unwrap_err().to_string();
174 assert!(err.contains("was cancelled"), "{err}");
175 }
176}