1use std::io;
2
3use arrow::datatypes::SchemaRef;
4use arrow::record_batch::RecordBatch;
5
6pub 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
13pub 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
54pub 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}