Skip to main content

rs_arrow_paste_batch/
lib.rs

1use std::io;
2
3use arrow::datatypes::SchemaRef;
4use arrow::record_batch::RecordBatch;
5
6/// Merges two schemas into a new schema.
7pub fn paste_schema(s0: SchemaRef, s1: SchemaRef) -> SchemaRef {
8    let mut fields = s0.fields().to_vec();
9    fields.extend_from_slice(s1.fields());
10    std::sync::Arc::new(arrow::datatypes::Schema::new(fields))
11}
12
13/// Pastes two iterators of record batches together.
14///
15/// If the schemas are not provided, they are inferred from the first batch of each iterator.
16///
17/// # Arguments
18///
19/// * `b0` - The first iterator of record batches.
20/// * `b1` - The second iterator of record batches.
21/// * `os0` - An optional schema for the first iterator.
22/// * `os1` - An optional schema for the second iterator.
23///
24/// # Returns
25///
26/// An iterator that yields the pasted record batches.
27pub fn paste_sync<I, J>(
28    mut b0: I,
29    mut b1: J,
30    os0: Option<SchemaRef>,
31    os1: Option<SchemaRef>,
32) -> Box<dyn Iterator<Item = Result<RecordBatch, io::Error>> + 'static>
33where
34    I: Iterator<Item = Result<RecordBatch, io::Error>> + 'static,
35    J: Iterator<Item = Result<RecordBatch, io::Error>> + 'static,
36{
37    let first_rb0 = b0.next();
38    let first_rb1 = b1.next();
39
40    let (s0, s1) = match (os0, os1) {
41        (Some(s0), Some(s1)) => (s0, s1),
42        _ => match (first_rb0.as_ref(), first_rb1.as_ref()) {
43            (Some(Ok(rb0)), Some(Ok(rb1))) => (rb0.schema(), rb1.schema()),
44            _ => return Box::new(std::iter::empty()),
45        },
46    };
47
48    let bz = first_rb0.into_iter().chain(b0);
49    let bo = first_rb1.into_iter().chain(b1);
50
51    Box::new(paste_sync_alt(bz, bo, s0, s1))
52}
53
54/// Pastes two iterators of record batches together, given the schemas.
55///
56/// # Arguments
57///
58/// * `b0` - The first iterator of record batches.
59/// * `b1` - The second iterator of record batches.
60/// * `s0` - The schema for the first iterator.
61/// * `s1` - The schema for the second iterator.
62///
63/// # Returns
64///
65/// An iterator that yields the pasted record batches.
66pub fn paste_sync_alt<I, J>(
67    b0: I,
68    b1: J,
69    s0: SchemaRef,
70    s1: SchemaRef,
71) -> impl Iterator<Item = Result<RecordBatch, io::Error>>
72where
73    I: Iterator<Item = Result<RecordBatch, io::Error>> + 'static,
74    J: Iterator<Item = Result<RecordBatch, io::Error>> + 'static,
75{
76    let schema = paste_schema(s0, s1);
77    b0.zip(b1).map(move |(rb0, rb1)| {
78        let rb0 = rb0?;
79        let rb1 = rb1?;
80        let mut columns = rb0.columns().to_vec();
81        columns.extend_from_slice(rb1.columns());
82        RecordBatch::try_new(schema.clone(), columns).map_err(io::Error::other)
83    })
84}
85
86#[cfg(test)]
87mod tests {
88    use super::*;
89    use arrow::datatypes::{DataType, Field, Schema};
90    use std::sync::Arc;
91
92    #[test]
93    fn test_paste_schema() {
94        let s0 = Arc::new(Schema::new(vec![Field::new("a", DataType::Int64, false)]));
95        let s1 = Arc::new(Schema::new(vec![Field::new("b", DataType::Utf8, true)]));
96
97        let pasted_schema = paste_schema(s0.clone(), s1.clone());
98
99        let expected_fields = vec![
100            Field::new("a", DataType::Int64, false),
101            Field::new("b", DataType::Utf8, true),
102        ];
103        let expected_schema = Arc::new(Schema::new(expected_fields));
104
105        assert_eq!(pasted_schema.fields().len(), 2);
106        assert_eq!(pasted_schema, expected_schema);
107    }
108
109    #[test]
110    fn test_paste_sync() -> Result<(), Box<dyn std::error::Error>> {
111        let s0 = Arc::new(Schema::new(vec![Field::new("a", DataType::Int64, false)]));
112        let s1 = Arc::new(Schema::new(vec![Field::new("b", DataType::Utf8, true)]));
113
114        let a = arrow::array::Int64Array::from(vec![1, 2, 3]);
115        let b = arrow::array::StringArray::from(vec!["a", "b", "c"]);
116
117        let rb0 = RecordBatch::try_new(s0.clone(), vec![Arc::new(a)])?;
118        let rb1 = RecordBatch::try_new(s1.clone(), vec![Arc::new(b)])?;
119
120        let b0 = vec![Ok(rb0)].into_iter();
121        let b1 = vec![Ok(rb1)].into_iter();
122
123        let mut pasted = paste_sync(b0, b1, None, None);
124        let pasted_rb = pasted.next().ok_or("no next value")??;
125
126        let expected_fields = vec![
127            Field::new("a", DataType::Int64, false),
128            Field::new("b", DataType::Utf8, true),
129        ];
130        let expected_schema = Arc::new(Schema::new(expected_fields));
131
132        assert_eq!(pasted_rb.schema(), expected_schema);
133        assert_eq!(pasted_rb.num_columns(), 2);
134        assert_eq!(pasted_rb.num_rows(), 3);
135        Ok(())
136    }
137
138    #[test]
139    fn test_paste_sync_alt() -> Result<(), Box<dyn std::error::Error>> {
140        let s0 = Arc::new(Schema::new(vec![Field::new("a", DataType::Int64, false)]));
141        let s1 = Arc::new(Schema::new(vec![Field::new("b", DataType::Utf8, true)]));
142
143        let a = arrow::array::Int64Array::from(vec![1, 2, 3]);
144        let b = arrow::array::StringArray::from(vec!["a", "b", "c"]);
145
146        let rb0 = RecordBatch::try_new(s0.clone(), vec![Arc::new(a)])?;
147        let rb1 = RecordBatch::try_new(s1.clone(), vec![Arc::new(b)])?;
148
149        let b0 = vec![Ok(rb0)].into_iter();
150        let b1 = vec![Ok(rb1)].into_iter();
151
152        let mut pasted = paste_sync_alt(b0, b1, s0, s1);
153        let pasted_rb = pasted.next().ok_or("no next value")??;
154
155        let expected_fields = vec![
156            Field::new("a", DataType::Int64, false),
157            Field::new("b", DataType::Utf8, true),
158        ];
159        let expected_schema = Arc::new(Schema::new(expected_fields));
160
161        assert_eq!(pasted_rb.schema(), expected_schema);
162        assert_eq!(pasted_rb.num_columns(), 2);
163        assert_eq!(pasted_rb.num_rows(), 3);
164        Ok(())
165    }
166}