polars_python/io/
arrow_c_stream.rs1use arrow::array::{Array, StructArray};
2use arrow::ffi::{ArrowArrayStream, ArrowArrayStreamReader};
3use parking_lot::Mutex;
4use polars::prelude::*;
5use pyo3::exceptions::PyValueError;
6use pyo3::prelude::*;
7
8use crate::conversion::Wrap;
9use crate::dataframe::PyDataFrame;
10use crate::error::PyPolarsErr;
11use crate::series::{call_arrow_c_stream, open_stream_capsule};
12
13struct ReaderState {
14 reader: ArrowArrayStreamReader<Box<ArrowArrayStream>>,
15 projection: Option<Option<PlHashSet<PlSmallStr>>>,
19}
20
21#[pyclass]
22pub struct PyArrowCStreamReader {
23 state: Mutex<ReaderState>,
24 schema: Schema,
25}
26
27#[pymethods]
28impl PyArrowCStreamReader {
29 #[new]
30 fn new(ob: &Bound<PyAny>) -> PyResult<Self> {
31 let capsule = call_arrow_c_stream(ob)?;
32 let reader = open_stream_capsule(&capsule)?;
33
34 let ArrowDataType::Struct(fields) = &reader.field().dtype else {
35 return Err(PyValueError::new_err(
36 "Arrow C Stream schema must be a struct type",
37 ));
38 };
39 let schema = Schema::from_iter(fields.iter().map(Field::from));
40
41 Ok(Self {
42 state: Mutex::new(ReaderState {
43 reader,
44 projection: None,
45 }),
46 schema,
47 })
48 }
49
50 #[getter]
51 fn schema(&self) -> Wrap<Schema> {
52 Wrap(self.schema.clone())
53 }
54
55 fn next_batch(&self, with_columns: Option<Vec<PlSmallStr>>) -> PyResult<Option<PyDataFrame>> {
56 let mut state = self.state.lock();
57 if state.projection.is_none() {
58 state.projection = Some(with_columns.map(|cols| cols.into_iter().collect()));
59 }
60
61 let array = match unsafe { state.reader.next() } {
62 Some(Ok(array)) => array,
63 Some(Err(e)) => return Err(PyPolarsErr::from(e).into()),
64 None => return Ok(None),
65 };
66
67 let projection = state.projection.as_ref().unwrap().as_ref();
68 let df = struct_array_to_df(array, projection).map_err(PyPolarsErr::from)?;
69 Ok(Some(PyDataFrame::new(df)))
70 }
71}
72
73fn struct_array_to_df(
74 array: Box<dyn Array>,
75 projection: Option<&PlHashSet<PlSmallStr>>,
76) -> PolarsResult<DataFrame> {
77 let struct_array = array.as_any().downcast_ref::<StructArray>().ok_or_else(
78 || polars_err!(ComputeError: "expected a StructArray from the Arrow C Stream"),
79 )?;
80
81 let columns = struct_array
82 .values()
83 .iter()
84 .zip(struct_array.fields())
85 .filter(|(_, field)| projection.is_none_or(|proj| proj.contains(&field.name)))
86 .map(|(arr, field)| unsafe {
87 Series::_try_from_arrow_unchecked(field.name.clone(), vec![arr.clone()], arr.dtype())
88 .map(Series::into_column)
89 })
90 .collect::<PolarsResult<Vec<_>>>()?;
91
92 DataFrame::new_infer_height(columns)
93}