Skip to main content

postg_arrow/
export.rs

1use crate::types::pg_oid_to_arrow;
2use arrow::array::{ArrayRef, Int16Builder, Int32Builder, Int64Builder, Float32Builder, Float64Builder, BooleanBuilder, LargeBinaryBuilder, LargeStringBuilder};
3use arrow::datatypes::{Field, Schema, DataType};
4use arrow::record_batch::RecordBatch;
5use bytes::{Buf, BytesMut};
6use futures::{Stream, StreamExt};
7use sqlx::{PgConnection, Executor, Column, Statement};
8use sqlx::postgres::PgTypeInfo;
9use anyhow::{anyhow, Result};
10use std::sync::Arc;
11use std::convert::TryInto;
12
13enum ColumnBuilder {
14    Int16(Int16Builder),
15    Int32(Int32Builder),
16    Int64(Int64Builder),
17    Float32(Float32Builder),
18    Float64(Float64Builder),
19    Boolean(BooleanBuilder),
20    LargeBinary(LargeBinaryBuilder),
21    LargeUtf8(LargeStringBuilder),
22}
23
24impl ColumnBuilder {
25    fn append_null(&mut self) {
26        match self {
27            ColumnBuilder::Int16(b) => b.append_null(),
28            ColumnBuilder::Int32(b) => b.append_null(),
29            ColumnBuilder::Int64(b) => b.append_null(),
30            ColumnBuilder::Float32(b) => b.append_null(),
31            ColumnBuilder::Float64(b) => b.append_null(),
32            ColumnBuilder::Boolean(b) => b.append_null(),
33            ColumnBuilder::LargeBinary(b) => b.append_null(),
34            ColumnBuilder::LargeUtf8(b) => b.append_null(),
35        }
36    }
37
38    fn finish(&mut self) -> ArrayRef {
39        match self {
40            ColumnBuilder::Int16(b) => Arc::new(b.finish()),
41            ColumnBuilder::Int32(b) => Arc::new(b.finish()),
42            ColumnBuilder::Int64(b) => Arc::new(b.finish()),
43            ColumnBuilder::Float32(b) => Arc::new(b.finish()),
44            ColumnBuilder::Float64(b) => Arc::new(b.finish()),
45            ColumnBuilder::Boolean(b) => Arc::new(b.finish()),
46            ColumnBuilder::LargeBinary(b) => Arc::new(b.finish()),
47            ColumnBuilder::LargeUtf8(b) => Arc::new(b.finish()),
48        }
49    }
50
51    fn new_from_datatype(dt: &DataType) -> Result<Self> {
52        match dt {
53            DataType::Int16 => Ok(ColumnBuilder::Int16(Int16Builder::with_capacity(1000))),
54            DataType::Int32 => Ok(ColumnBuilder::Int32(Int32Builder::with_capacity(1000))),
55            DataType::Int64 => Ok(ColumnBuilder::Int64(Int64Builder::with_capacity(1000))),
56            DataType::Float32 => Ok(ColumnBuilder::Float32(Float32Builder::with_capacity(1000))),
57            DataType::Float64 => Ok(ColumnBuilder::Float64(Float64Builder::with_capacity(1000))),
58            DataType::Boolean => Ok(ColumnBuilder::Boolean(BooleanBuilder::with_capacity(1000))),
59            DataType::LargeBinary => Ok(ColumnBuilder::LargeBinary(LargeBinaryBuilder::with_capacity(1000, 1024))),
60            DataType::LargeUtf8 => Ok(ColumnBuilder::LargeUtf8(LargeStringBuilder::with_capacity(1000, 1024))),
61            _ => Err(anyhow!("Unsupported type in export: {:?}", dt)),
62        }
63    }
64}
65
66pub async fn query_to_arrow<'a>(
67    conn: &'a mut PgConnection,
68    query: &str,
69) -> Result<(Arc<Schema>, impl Stream<Item = Result<RecordBatch>> + 'a)> {
70    let describe_query = format!("SELECT * FROM ({}) AS _t LIMIT 0", query);
71    let stmt = conn.prepare(describe_query.as_str()).await?;
72
73    let mut fields: Vec<Field> = Vec::new();
74    for col in stmt.columns() {
75        let type_info: &PgTypeInfo = col.type_info();
76        let oid: u32 = type_info.oid().map(|o| o.0).unwrap_or(0);
77        let arrow_type = pg_oid_to_arrow(oid)?;
78        fields.push(Field::new(col.name(), arrow_type, true));
79    }
80    let schema = Arc::new(Schema::new(fields));
81
82    let copy_query = format!("COPY ({}) TO STDOUT WITH (FORMAT binary)", query);
83    let mut copy_out = conn.copy_out_raw(copy_query.as_str()).await?;
84
85    let schema_clone = schema.clone();
86    
87    let stream = async_stream::try_stream! {
88        let mut buf = BytesMut::new();
89        let mut header_parsed = false;
90
91        let mut builders: Vec<ColumnBuilder> = Vec::new();
92        for f in schema_clone.fields() {
93            builders.push(ColumnBuilder::new_from_datatype(f.data_type())?);
94        }
95
96        let mut row_count = 0;
97
98        while let Some(chunk_res) = copy_out.next().await {
99            let chunk: bytes::Bytes = chunk_res?;
100            buf.extend_from_slice(&chunk);
101
102            if !header_parsed {
103                if buf.len() >= 19 {
104                    buf.advance(11); // Signature
105                    let _flags = buf.get_i32();
106                    let ext_len = buf.get_i32();
107                    buf.advance(ext_len as usize);
108                    header_parsed = true;
109                } else {
110                    continue;
111                }
112            }
113
114            loop {
115                if buf.len() < 2 {
116                    break;
117                }
118                let num_cols = i16::from_be_bytes([buf[0], buf[1]]);
119                if num_cols == -1 {
120                    buf.advance(2); // Trailer
121                    break;
122                }
123
124                let mut offset = 2;
125                let mut can_read = true;
126                for _ in 0..num_cols {
127                    if offset + 4 > buf.len() {
128                        can_read = false;
129                        break;
130                    }
131                    let col_len = i32::from_be_bytes([buf[offset], buf[offset+1], buf[offset+2], buf[offset+3]]);
132                    offset += 4;
133                    if col_len > 0 {
134                        if offset + col_len as usize > buf.len() {
135                            can_read = false;
136                            break;
137                        }
138                        offset += col_len as usize;
139                    }
140                }
141
142                if !can_read {
143                    break;
144                }
145
146                buf.advance(2); // num_cols
147                for i in 0..num_cols as usize {
148                    let col_len = buf.get_i32();
149                    if col_len == -1 {
150                        builders[i].append_null();
151                    } else {
152                        let data = buf.split_to(col_len as usize);
153                        match &mut builders[i] {
154                            ColumnBuilder::Int16(b) => {
155                                b.append_value(i16::from_be_bytes(data[..2].try_into().unwrap()));
156                            }
157                            ColumnBuilder::Int32(b) => {
158                                b.append_value(i32::from_be_bytes(data[..4].try_into().unwrap()));
159                            }
160                            ColumnBuilder::Int64(b) => {
161                                b.append_value(i64::from_be_bytes(data[..8].try_into().unwrap()));
162                            }
163                            ColumnBuilder::Float32(b) => {
164                                b.append_value(f32::from_be_bytes(data[..4].try_into().unwrap()));
165                            }
166                            ColumnBuilder::Float64(b) => {
167                                b.append_value(f64::from_be_bytes(data[..8].try_into().unwrap()));
168                            }
169                            ColumnBuilder::Boolean(b) => {
170                                b.append_value(data[0] != 0);
171                            }
172                            ColumnBuilder::LargeBinary(b) => {
173                                b.append_value(&data);
174                            }
175                            ColumnBuilder::LargeUtf8(b) => {
176                                let val = std::str::from_utf8(&data)?;
177                                b.append_value(val);
178                            }
179                        }
180                    }
181                }
182                
183                row_count += 1;
184                if row_count >= 1000 {
185                    let mut arrays: Vec<ArrayRef> = Vec::new();
186                    for i in 0..num_cols as usize {
187                        arrays.push(builders[i].finish());
188                        builders[i] = ColumnBuilder::new_from_datatype(schema_clone.field(i).data_type())?;
189                    }
190                    let batch = RecordBatch::try_new(schema_clone.clone(), arrays)?;
191                    yield batch;
192                    row_count = 0;
193                }
194            }
195        }
196
197        if row_count > 0 {
198            let mut arrays: Vec<ArrayRef> = Vec::new();
199            for i in 0..schema_clone.fields().len() {
200                arrays.push(builders[i].finish());
201            }
202            let batch = RecordBatch::try_new(schema_clone.clone(), arrays)?;
203            yield batch;
204        }
205    };
206
207    Ok((schema, stream))
208}