Skip to main content

connectorx/destinations/arrowstream/
mod.rs

1//! Destination implementation for Arrow and Polars.
2
3mod arrow_assoc;
4mod errors;
5mod funcs;
6pub mod typesystem;
7
8pub use self::errors::{ArrowDestinationError, Result};
9pub use self::typesystem::ArrowTypeSystem;
10use super::{Consume, Destination, DestinationPartition};
11use crate::constants::RECORD_BATCH_SIZE;
12use crate::data_order::DataOrder;
13use crate::typesystem::{Realize, TypeAssoc, TypeSystem};
14use anyhow::anyhow;
15use arrow::{datatypes::Schema, record_batch::RecordBatch};
16use arrow_assoc::ArrowAssoc;
17use fehler::{throw, throws};
18use funcs::{FFinishBuilder, FNewBuilder, FNewField};
19use itertools::Itertools;
20use std::{
21    any::Any,
22    sync::{
23        mpsc::{channel, Receiver, Sender},
24        Arc,
25    },
26};
27
28type Builder = Box<dyn Any + Send>;
29type Builders = Vec<Builder>;
30
31pub struct ArrowDestination {
32    schema: Vec<ArrowTypeSystem>,
33    names: Vec<String>,
34    arrow_schema: Arc<Schema>,
35    batch_size: usize,
36    sender: Option<Sender<RecordBatch>>,
37    receiver: Receiver<RecordBatch>,
38}
39
40impl Default for ArrowDestination {
41    fn default() -> Self {
42        let (tx, rx) = channel();
43        ArrowDestination {
44            schema: vec![],
45            names: vec![],
46            arrow_schema: Arc::new(Schema::empty()),
47            batch_size: RECORD_BATCH_SIZE,
48            sender: Some(tx),
49            receiver: rx,
50        }
51    }
52}
53
54impl ArrowDestination {
55    pub fn new() -> Self {
56        Self::default()
57    }
58
59    pub fn new_with_batch_size(batch_size: usize) -> Self {
60        let (tx, rx) = channel();
61        ArrowDestination {
62            schema: vec![],
63            names: vec![],
64            arrow_schema: Arc::new(Schema::empty()),
65            batch_size,
66            sender: Some(tx),
67            receiver: rx,
68        }
69    }
70}
71
72impl Destination for ArrowDestination {
73    const DATA_ORDERS: &'static [DataOrder] = &[DataOrder::ColumnMajor, DataOrder::RowMajor];
74    type TypeSystem = ArrowTypeSystem;
75    type Partition<'a> = ArrowPartitionWriter;
76    type Error = ArrowDestinationError;
77
78    fn needs_count(&self) -> bool {
79        false
80    }
81
82    #[throws(ArrowDestinationError)]
83    fn allocate<S: AsRef<str>>(
84        &mut self,
85        _nrow: usize,
86        names: &[S],
87        schema: &[ArrowTypeSystem],
88        data_order: DataOrder,
89    ) {
90        // todo: support colmajor
91        if !matches!(data_order, DataOrder::RowMajor) {
92            throw!(crate::errors::ConnectorXError::UnsupportedDataOrder(
93                data_order
94            ))
95        }
96
97        // parse the metadata
98        self.schema = schema.to_vec();
99        self.names = names.iter().map(|n| n.as_ref().to_string()).collect();
100        let fields = self
101            .schema
102            .iter()
103            .zip_eq(&self.names)
104            .map(|(&dt, h)| Ok(Realize::<FNewField>::realize(dt)?(h.as_str())))
105            .collect::<Result<Vec<_>>>()?;
106        self.arrow_schema = Arc::new(Schema::new(fields));
107    }
108
109    #[throws(ArrowDestinationError)]
110    fn partition(&mut self, counts: usize) -> Vec<Self::Partition<'_>> {
111        let mut partitions = vec![];
112        let sender = self.sender.take().unwrap();
113        for _ in 0..counts {
114            partitions.push(ArrowPartitionWriter::new(
115                self.schema.clone(),
116                Arc::clone(&self.arrow_schema),
117                self.batch_size,
118                sender.clone(),
119            )?);
120        }
121        partitions
122        // self.sender should be freed
123    }
124
125    fn schema(&self) -> &[ArrowTypeSystem] {
126        self.schema.as_slice()
127    }
128}
129
130impl ArrowDestination {
131    #[throws(ArrowDestinationError)]
132    pub fn arrow(self) -> Vec<RecordBatch> {
133        if self.sender.is_some() {
134            // should not happen since it is dropped after partition
135            // but need to make sure here otherwise recv will be blocked forever
136            std::mem::drop(self.sender);
137        }
138        let mut data = vec![];
139        while let Ok(rb) = self.receiver.recv() {
140            data.push(rb);
141        }
142        data
143    }
144
145    #[throws(ArrowDestinationError)]
146    pub fn record_batch(&mut self) -> Option<RecordBatch> {
147        self.receiver.recv().ok()
148    }
149
150    pub fn empty_batch(&self) -> RecordBatch {
151        RecordBatch::new_empty(self.arrow_schema.clone())
152    }
153
154    pub fn arrow_schema(&self) -> Arc<Schema> {
155        self.arrow_schema.clone()
156    }
157
158    pub fn names(&self) -> &[String] {
159        self.names.as_slice()
160    }
161}
162
163pub struct ArrowPartitionWriter {
164    schema: Vec<ArrowTypeSystem>,
165    builders: Option<Builders>,
166    current_row: usize,
167    current_col: usize,
168    arrow_schema: Arc<Schema>,
169    batch_size: usize,
170    sender: Option<Sender<RecordBatch>>,
171}
172
173// unsafe impl Sync for ArrowPartitionWriter {}
174
175impl ArrowPartitionWriter {
176    #[throws(ArrowDestinationError)]
177    fn new(
178        schema: Vec<ArrowTypeSystem>,
179        arrow_schema: Arc<Schema>,
180        batch_size: usize,
181        sender: Sender<RecordBatch>,
182    ) -> Self {
183        let mut pw = ArrowPartitionWriter {
184            schema,
185            builders: None,
186            current_row: 0,
187            current_col: 0,
188            arrow_schema,
189            batch_size,
190            sender: Some(sender),
191        };
192        pw.allocate()?;
193        pw
194    }
195
196    #[throws(ArrowDestinationError)]
197    fn allocate(&mut self) {
198        let builders = self
199            .schema
200            .iter()
201            .map(|dt| Ok(Realize::<FNewBuilder>::realize(*dt)?(self.batch_size)))
202            .collect::<Result<Vec<_>>>()?;
203        self.builders.replace(builders);
204    }
205
206    #[throws(ArrowDestinationError)]
207    fn flush(&mut self) {
208        let builders = self
209            .builders
210            .take()
211            .unwrap_or_else(|| panic!("arrow builder is none when flush!"));
212        let columns = builders
213            .into_iter()
214            .zip(self.schema.iter())
215            .map(|(builder, &dt)| Realize::<FFinishBuilder>::realize(dt)?(builder))
216            .collect::<std::result::Result<Vec<_>, crate::errors::ConnectorXError>>()?;
217        let rb = RecordBatch::try_new(Arc::clone(&self.arrow_schema), columns)?;
218        self.sender.as_ref().and_then(|s| s.send(rb).ok());
219
220        self.current_row = 0;
221        self.current_col = 0;
222    }
223}
224
225impl<'a> DestinationPartition<'a> for ArrowPartitionWriter {
226    type TypeSystem = ArrowTypeSystem;
227    type Error = ArrowDestinationError;
228
229    #[throws(ArrowDestinationError)]
230    fn finalize(&mut self) {
231        if self.builders.is_some() {
232            self.flush()?;
233        }
234        // need to release the sender so receiver knows when the stream is exhasted
235        std::mem::drop(self.sender.take());
236    }
237
238    #[throws(ArrowDestinationError)]
239    fn aquire_row(&mut self, _n: usize) -> usize {
240        self.current_row
241    }
242
243    fn ncols(&self) -> usize {
244        self.schema.len()
245    }
246}
247
248impl<'a, T> Consume<T> for ArrowPartitionWriter
249where
250    T: TypeAssoc<<Self as DestinationPartition<'a>>::TypeSystem> + ArrowAssoc + 'static,
251{
252    type Error = ArrowDestinationError;
253
254    #[throws(ArrowDestinationError)]
255    fn consume(&mut self, value: T) {
256        let col = self.current_col;
257        self.current_col = (self.current_col + 1) % self.ncols();
258        self.schema[col].check::<T>()?;
259
260        loop {
261            match &mut self.builders {
262                Some(builders) => {
263                    <T as ArrowAssoc>::append(
264                        builders[col]
265                            .downcast_mut::<T::Builder>()
266                            .ok_or_else(|| anyhow!("cannot cast arrow builder for append"))?,
267                        value,
268                    )?;
269                    break;
270                }
271                None => self.allocate()?, // allocate if builders are not initialized
272            }
273        }
274
275        // flush if exceed batch_size
276        if self.current_col == 0 {
277            self.current_row += 1;
278            if self.current_row >= self.batch_size {
279                self.flush()?;
280                self.allocate()?;
281            }
282        }
283    }
284}