Skip to main content

polars_python/io/
arrow_c_stream.rs

1use 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    // Resolved from the first `next_batch` call's `with_columns` and reused
16    // after: the engine calls `next_batch` with the same projection on every
17    // batch of a given scan, so re-parsing it per batch would be wasted work.
18    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}