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); 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); 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); 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}