Skip to main content

questdb/egress/arrow/
polars.rs

1//! Polars sub-feature: `RecordBatch ↔ DataFrame` via Arrow C Data Interface.
2
3use std::sync::Arc;
4
5use arrow::array::ArrayData;
6use arrow::array::types::UInt32Type;
7use arrow::array::{Array, ArrayRef, DictionaryArray, RecordBatch};
8use arrow::datatypes::{DataType as ArrowDataType, SchemaRef};
9use polars::frame::DataFrame;
10use polars::prelude::{
11    Categorical32Type, CategoricalChunked, CategoricalMapping, CategoricalPhysical, Categories,
12    Column, DataType as PlDataType, IDX_DTYPE, IdxCa, IntoColumn, IntoSeries, PlSmallStr, Series,
13};
14
15use crate::egress::Cursor;
16use crate::egress::arrow::has_tentative_array;
17use crate::egress::symbol_dict::SymbolDict;
18use crate::error::{Error, ErrorCode, Result, fmt};
19
20// FFI cross-crate helpers in `crate::ingress::polars`.
21
22impl Cursor<'_> {
23    /// Decode one batch as a Polars [`DataFrame`]. `Ok(None)` on
24    /// stream end.
25    ///
26    /// This is the low-level per-batch entry point and does **not**
27    /// detect mid-stream Arrow schema drift; if a later batch's
28    /// schema differs from earlier ones the resulting DataFrames will
29    /// simply disagree on columns. Use
30    /// [`Cursor::iter_polars`](Cursor::iter_polars)
31    /// for a drift-checked iterator, or
32    /// [`Cursor::fetch_all_polars`] / [`Cursor::as_arrow_reader`]
33    /// for higher-level adapters that pin the schema on first batch.
34    pub fn next_polars(&mut self) -> Result<Option<DataFrame>> {
35        match self.next_arrow_batch_inner(None, false)? {
36            None => Ok(None),
37            Some(rb) => Ok(Some(self.batch_to_dataframe(rb)?)),
38        }
39    }
40
41    /// Per-batch `RecordBatch → DataFrame` via the cursor's persistent
42    /// [`SymbolRegistry`], instead of rebuilding the categorical mapping every
43    /// batch (the high-cardinality collapse).
44    fn batch_to_dataframe(&mut self, rb: RecordBatch) -> Result<DataFrame> {
45        let modes = self.symbol_delta_modes().to_vec();
46        let registry = self.symbol_registry_synced()?;
47        build_dataframe(rb, &modes, registry)
48    }
49
50    /// Eagerly drain into one chunked Polars [`DataFrame`]. A stream
51    /// that yields a schema but no batches becomes an empty DataFrame;
52    /// only a stream without a schema (e.g. cancelled pre-prelude)
53    /// errors as `NoSchema`. Drift detection is inherited from
54    /// [`Cursor::iter_polars`].
55    pub fn fetch_all_polars(&mut self) -> Result<DataFrame> {
56        // Materialise-whole: the full result is built internally before
57        // anything is handed back, so a mid-query failover replays the
58        // query from `batch_seq 0` transparently. Opt into replay and
59        // drop the partial accumulation when the cursor reports a reset.
60        self.enable_internal_replay();
61        let mut iter = self.iter_polars()?;
62        let mut resets_seen = iter.failover_resets();
63        let mut acc: Option<DataFrame> = None;
64        loop {
65            // Manual drive (not `for`/`by_ref`) so the reset counter can be
66            // polled between batches without holding an iterator borrow.
67            let Some(item) = iter.next() else { break };
68            let df = item?;
69            let resets_now = iter.failover_resets();
70            if resets_now != resets_seen {
71                resets_seen = resets_now;
72                acc = None;
73            }
74            acc = Some(match acc {
75                None => df,
76                Some(mut prev) => {
77                    if prev.height() == 0 && prev.schema() != df.schema() {
78                        df
79                    } else {
80                        prev.vstack_mut_owned(df)
81                            .map_err(|e| fmt!(ArrowExport, "polars vstack failed: {}", e))?;
82                        prev
83                    }
84                }
85            });
86        }
87        let schema = iter.schema();
88        match acc {
89            Some(df) => Ok(df),
90            None => record_batch_to_dataframe(RecordBatch::new_empty(schema)),
91        }
92    }
93}
94
95/// Drift-checked iterator yielding Polars [`DataFrame`]s, one per
96/// QWP batch. Built by [`Cursor::iter_polars`]. Snapshots the first
97/// batch's Arrow schema at construction and poisons (terminates) on
98/// mid-stream schema drift.
99pub struct CursorPolarsIter<'r, 'c> {
100    cursor: &'c mut Cursor<'r>,
101    schema: SchemaRef,
102    pending: Option<RecordBatch>,
103    poisoned: bool,
104    /// `Cursor::failover_resets()` at the point `schema` was pinned. A
105    /// later batch arriving with a higher count is the first frame of a
106    /// transparently-replayed query (re-read from `batch_seq 0` on a new
107    /// endpoint), so the pinned schema is re-snapshotted from it rather
108    /// than treated as drift. `fetch_all_polars` reads the same counter to
109    /// discard its partial `vstack` accumulation.
110    resets_at_pin: u32,
111}
112
113impl<'r, 'c> CursorPolarsIter<'r, 'c> {
114    pub(crate) fn new(cursor: &'c mut Cursor<'r>) -> Result<Self> {
115        let first = cursor.next_arrow_batch_inner(None, false)?.ok_or_else(|| {
116            Error::new(
117                ErrorCode::NoSchema,
118                "no batch produced; nothing to snapshot",
119            )
120        })?;
121        let schema = first.schema();
122        let resets_at_pin = cursor.failover_resets();
123        Ok(Self {
124            cursor,
125            schema,
126            pending: Some(first),
127            poisoned: false,
128            resets_at_pin,
129        })
130    }
131
132    /// First batch's schema. Upgrades on tentative→firm ndim
133    /// (see [`has_tentative_array`]).
134    pub fn schema(&self) -> SchemaRef {
135        self.schema.clone()
136    }
137
138    /// Reconnect count observed by the underlying cursor. `fetch_all_polars`
139    /// polls this between batches: an increase means the query was replayed
140    /// from scratch, so anything accumulated so far must be dropped.
141    pub(crate) fn failover_resets(&self) -> u32 {
142        self.cursor.failover_resets()
143    }
144}
145
146impl Iterator for CursorPolarsIter<'_, '_> {
147    type Item = Result<DataFrame>;
148
149    fn next(&mut self) -> Option<Self::Item> {
150        if self.poisoned {
151            return None;
152        }
153        let rb = if let Some(rb) = self.pending.take() {
154            rb
155        } else {
156            // A transparent mid-query failover re-reads the result from
157            // `batch_seq 0` on a new endpoint. Pass `None` (no drift check)
158            // for that frame so the new node's batch 0 isn't rejected, then
159            // require it to match the pinned schema so the iterator's
160            // `schema()` stays stable across the replay.
161            let drift_check = if self.cursor.failover_resets() == self.resets_at_pin {
162                Some(&self.schema)
163            } else {
164                None
165            };
166            match self.cursor.next_arrow_batch_inner(drift_check, false) {
167                Ok(Some(rb)) => {
168                    if self.cursor.failover_resets() != self.resets_at_pin {
169                        if rb.schema() != self.schema {
170                            self.poisoned = true;
171                            return Some(Err(Error::new(
172                                ErrorCode::SchemaDrift,
173                                "post-failover replay returned a different \
174                                 schema; the iterator pins the first batch's \
175                                 schema. Use Cursor::next_polars to handle \
176                                 drift explicitly",
177                            )));
178                        }
179                        self.resets_at_pin = self.cursor.failover_resets();
180                    } else if has_tentative_array(&self.schema) && rb.schema() != self.schema {
181                        self.poisoned = true;
182                        return Some(Err(Error::new(
183                            ErrorCode::SchemaDrift,
184                            "tentative→firm ndim upgrade mid-stream; the \
185                             iterator pins the first batch's schema. Use \
186                             Cursor::next_polars to handle drift explicitly",
187                        )));
188                    }
189                    rb
190                }
191                Ok(None) => {
192                    self.poisoned = true;
193                    return None;
194                }
195                Err(e) => {
196                    self.poisoned = true;
197                    return Some(Err(e));
198                }
199            }
200        };
201        let df = self.cursor.batch_to_dataframe(rb);
202        if df.is_err() {
203            self.poisoned = true;
204        }
205        Some(df)
206    }
207}
208
209/// [`RecordBatch`] → Polars [`DataFrame`] via the Arrow C Data Interface.
210/// Zero-copy for primitive/string/binary; SYMBOL columns become a polars
211/// `Categorical`, using the same dictionary conversion as the cursor APIs.
212/// [`ErrorCode::ArrowExport`] on handoff failure.
213pub fn record_batch_to_dataframe(rb: RecordBatch) -> Result<DataFrame> {
214    let schema = rb.schema();
215    let mut columns: Vec<Column> = Vec::with_capacity(rb.num_columns());
216    // This standalone entry point has no per-cursor state to scope a
217    // `Categories` to, so SYMBOL strings intern into the process-global
218    // `Categories::global()` and are retained for the process lifetime. The
219    // cursor-driven paths (`next_polars` / `iter_polars` / `fetch_all_polars`)
220    // instead route through `build_dataframe`, which scopes interning to the
221    // cursor's `SymbolRegistry` (reset on CACHE_RESET). Using `global()` here
222    // also keeps the result's Categoricals compatible with any global
223    // Categoricals the caller already holds.
224    let cats = Categories::global();
225    let mapping = cats.mapping();
226    let cat_dtype = PlDataType::Categorical(cats, mapping);
227    for (col, field) in rb.columns().iter().zip(schema.fields().iter()) {
228        let name = field.name().as_str();
229        let series = if matches!(col.data_type(), ArrowDataType::Dictionary(_, _)) {
230            dictionary_to_categorical(name, col, &cat_dtype)?
231        } else {
232            import_polars_series(name, &col.to_data())?
233        };
234        columns.push(series.into_column());
235    }
236    crate::polars_ffi::df_from_columns(columns)
237        .map_err(|e| fmt!(ArrowExport, "DataFrame::new failed: {}", e))
238}
239
240fn import_polars_series(name: &str, array_data: &ArrayData) -> Result<Series> {
241    let (rs_array, rs_schema) = arrow::ffi::to_ffi(array_data)
242        .map_err(|e| fmt!(ArrowExport, "to_ffi failed for column '{}': {}", name, e))?;
243    let pa_schema = unsafe { crate::polars_ffi::rs_schema_into_pa(rs_schema) };
244    let pa_array = unsafe { crate::polars_ffi::rs_array_into_pa(rs_array) };
245    let pa_field = unsafe { polars_arrow::ffi::import_field_from_c(&pa_schema) }
246        .map_err(|e| fmt!(ArrowExport, "import_field_from_c('{}'): {}", name, e))?;
247    let pa_array_box = unsafe { polars_arrow::ffi::import_array_from_c(pa_array, pa_field.dtype) }
248        .map_err(|e| fmt!(ArrowExport, "import_array_from_c('{}'): {}", name, e))?;
249    Series::from_arrow(name.into(), pa_array_box)
250        .map_err(|e| fmt!(ArrowExport, "Series::from_arrow('{}'): {}", name, e))
251}
252
253/// Build a SYMBOL column's polars `Categorical` from its codes + dictionary
254/// (cast the small dictionary, then `take` by code), avoiding
255/// `Series::from_arrow`'s per-row remap. `Categories` exists since polars 0.50.
256///
257/// `cat_dtype` is the `Categorical` dtype to intern the dictionary strings into.
258/// All RESULT_BATCH frames of one query must pass the SAME dtype: polars
259/// vstack/concat — used by `fetch_all_polars` / `iter_polars` to stitch the
260/// per-batch frames — only accepts Categoricals that share one `Categories`
261/// identity, otherwise it errors "Categories name mismatch".
262fn dictionary_to_categorical(name: &str, col: &ArrayRef, cat_dtype: &PlDataType) -> Result<Series> {
263    let dict = col
264        .as_any()
265        .downcast_ref::<DictionaryArray<UInt32Type>>()
266        .ok_or_else(|| {
267            fmt!(
268                ArrowExport,
269                "SYMBOL '{}' is not Dictionary(UInt32, _)",
270                name
271            )
272        })?;
273
274    let values = import_polars_series(name, &dict.values().to_data())?;
275    let cat_dict = values.cast(cat_dtype).map_err(|e| {
276        fmt!(
277            ArrowExport,
278            "cast SYMBOL '{}' dict to Categorical: {}",
279            name,
280            e
281        )
282    })?;
283
284    let keys = import_polars_series(name, &dict.keys().to_data())?;
285    let idx: IdxCa = keys
286        .cast(&IDX_DTYPE)
287        .map_err(|e| fmt!(ArrowExport, "cast SYMBOL '{}' codes to index: {}", name, e))?
288        .idx()
289        .map_err(|e| {
290            fmt!(
291                ArrowExport,
292                "SYMBOL '{}' codes not an index dtype: {}",
293                name,
294                e
295            )
296        })?
297        .clone();
298    cat_dict
299        .take(&idx)
300        .map_err(|e| fmt!(ArrowExport, "gather SYMBOL '{}' codes: {}", name, e))
301}
302
303/// Per-cursor registry that interns a query's connection SYMBOL dictionary into
304/// one persistent polars `Categories` in QWP code order — so a global QWP code
305/// is its own physical categorical code and keys wrap straight into a
306/// Categorical, with no per-batch cast or gather.
307pub(crate) struct SymbolRegistry {
308    dtype: PlDataType,
309    mapping: Arc<CategoricalMapping>,
310    registered: usize,
311    /// Separate `Categories` for column-local (non-delta) SYMBOL columns routed
312    /// through the fallback. It must NOT share `mapping`: that one's physical id
313    /// == QWP global code by construction (`sync` interns in code order), and
314    /// interning arbitrary column-local strings into it would break that
315    /// alignment. A dedicated per-cursor `Categories` keeps the fallback's
316    /// retention scoped to the cursor (reset on CACHE_RESET with the rest of the
317    /// registry) instead of leaking into `Categories::global()` for the whole
318    /// process, while still sharing one identity across a stream's batches so
319    /// `fetch_all_polars` / `iter_polars` can vstack them.
320    fallback_dtype: PlDataType,
321}
322
323impl SymbolRegistry {
324    pub(crate) fn new() -> Self {
325        let cats = Categories::random(PlSmallStr::from("questdb_symbol"), CategoricalPhysical::U32);
326        let mapping = cats.mapping();
327        let dtype = PlDataType::Categorical(cats, mapping.clone());
328        let fallback_cats = Categories::random(
329            PlSmallStr::from("questdb_symbol_local"),
330            CategoricalPhysical::U32,
331        );
332        let fallback_mapping = fallback_cats.mapping();
333        let fallback_dtype = PlDataType::Categorical(fallback_cats, fallback_mapping);
334        Self {
335            dtype,
336            mapping,
337            registered: 0,
338            fallback_dtype,
339        }
340    }
341
342    fn local_dict_to_categorical(&self, name: &str, col: &ArrayRef) -> Result<Series> {
343        dictionary_to_categorical(name, col, &self.fallback_dtype)
344    }
345
346    pub(crate) fn sync(&mut self, dict: &SymbolDict) -> Result<()> {
347        // A CACHE_RESET clears the connection dict; the shrink restarts us.
348        if dict.len() < self.registered {
349            *self = Self::new();
350        }
351        for code in self.registered..dict.len() {
352            let s = dict.get(code as u32).ok_or_else(|| {
353                fmt!(
354                    ArrowExport,
355                    "symbol code {} missing from dict during registry sync",
356                    code
357                )
358            })?;
359            // `insert_cat` appends in call order, so the assigned id == `code`.
360            self.mapping
361                .insert_cat(s)
362                .map_err(|e| fmt!(ArrowExport, "register SYMBOL '{}': {}", s, e))?;
363        }
364        self.registered = dict.len();
365        Ok(())
366    }
367
368    fn categorical_from_keys(&self, name: &str, col: &ArrayRef) -> Result<Series> {
369        let dict = col
370            .as_any()
371            .downcast_ref::<DictionaryArray<UInt32Type>>()
372            .ok_or_else(|| {
373                fmt!(
374                    ArrowExport,
375                    "SYMBOL '{}' is not Dictionary(UInt32, _)",
376                    name
377                )
378            })?;
379        let keys = import_polars_series(name, &dict.keys().to_data())?;
380        let phys = keys
381            .u32()
382            .map_err(|e| fmt!(ArrowExport, "SYMBOL '{}' keys not u32: {}", name, e))?
383            .clone();
384        // SAFETY: `sync` registered every dict entry in code order, so each key is
385        // a valid physical code; the decoder bounds-checks every non-null key
386        // against `dict.len()`, and null rows (key 0) are masked by the imported
387        // key buffer's null bitmap.
388        let cat = unsafe {
389            CategoricalChunked::<Categorical32Type>::from_cats_and_dtype_unchecked(
390                phys,
391                self.dtype.clone(),
392            )
393        };
394        Ok(cat.into_series())
395    }
396}
397
398/// `RecordBatch → DataFrame` for the cursor-driven polars paths. Delta-mode
399/// SYMBOL columns (`delta_modes[i]`) resolve through the persistent `registry`;
400/// column-local SYMBOL columns fall back to
401/// [`SymbolRegistry::local_dict_to_categorical`], which scopes interning to the
402/// cursor's own `Categories` rather than the process-global one.
403fn build_dataframe(
404    rb: RecordBatch,
405    delta_modes: &[bool],
406    registry: &SymbolRegistry,
407) -> Result<DataFrame> {
408    let schema = rb.schema();
409    let mut columns: Vec<Column> = Vec::with_capacity(rb.num_columns());
410    for (i, (col, field)) in rb.columns().iter().zip(schema.fields().iter()).enumerate() {
411        let name = field.name().as_str();
412        let series = if matches!(col.data_type(), ArrowDataType::Dictionary(_, _)) {
413            if delta_modes.get(i).copied().unwrap_or(false) {
414                registry.categorical_from_keys(name, col)?
415            } else {
416                registry.local_dict_to_categorical(name, col)?
417            }
418        } else {
419            import_polars_series(name, &col.to_data())?
420        };
421        columns.push(series.into_column());
422    }
423    crate::polars_ffi::df_from_columns(columns)
424        .map_err(|e| fmt!(ArrowExport, "DataFrame::new failed: {}", e))
425}
426
427#[cfg(test)]
428mod tests {
429    use super::*;
430    use std::sync::Arc;
431
432    use arrow::array::builder::{Float64Builder, Int64Builder, StringBuilder};
433    use arrow::array::{ArrayRef, RecordBatch};
434    use arrow::datatypes::{DataType, Field, Schema as ArrowSchema};
435
436    fn rb_mixed() -> RecordBatch {
437        let mut ii = Int64Builder::new();
438        ii.append_value(1);
439        ii.append_value(2);
440        ii.append_value(3);
441        let mut ff = Float64Builder::new();
442        ff.append_value(1.5);
443        ff.append_value(2.5);
444        ff.append_value(3.5);
445        let mut ss = StringBuilder::new();
446        ss.append_value("a");
447        ss.append_value("b");
448        ss.append_value("c");
449        let schema = Arc::new(ArrowSchema::new(vec![
450            Field::new("i", DataType::Int64, false),
451            Field::new("f", DataType::Float64, false),
452            Field::new("s", DataType::Utf8, false),
453        ]));
454        RecordBatch::try_new(
455            schema,
456            vec![
457                Arc::new(ii.finish()) as ArrayRef,
458                Arc::new(ff.finish()) as ArrayRef,
459                Arc::new(ss.finish()) as ArrayRef,
460            ],
461        )
462        .unwrap()
463    }
464
465    #[test]
466    fn record_batch_to_dataframe_preserves_column_count_and_height() {
467        let rb = rb_mixed();
468        let df = record_batch_to_dataframe(rb).unwrap();
469        assert_eq!(df.width(), 3);
470        assert_eq!(df.height(), 3);
471        assert_eq!(df.select_at_idx(0).unwrap().name().as_str(), "i");
472        assert_eq!(df.select_at_idx(1).unwrap().name().as_str(), "f");
473        assert_eq!(df.select_at_idx(2).unwrap().name().as_str(), "s");
474    }
475
476    #[test]
477    fn record_batch_to_dataframe_preserves_int_values() {
478        let rb = rb_mixed();
479        let df = record_batch_to_dataframe(rb).unwrap();
480        let col = df.select_at_idx(0).unwrap();
481        let series = col.as_materialized_series();
482        let i64s = series.i64().unwrap();
483        assert_eq!(i64s.get(0), Some(1));
484        assert_eq!(i64s.get(1), Some(2));
485        assert_eq!(i64s.get(2), Some(3));
486    }
487
488    #[test]
489    fn record_batch_to_dataframe_preserves_string_values() {
490        let rb = rb_mixed();
491        let df = record_batch_to_dataframe(rb).unwrap();
492        let col = df.select_at_idx(2).unwrap();
493        let series = col.as_materialized_series();
494        let s = series.str().unwrap();
495        assert_eq!(s.get(0), Some("a"));
496        assert_eq!(s.get(1), Some("b"));
497        assert_eq!(s.get(2), Some("c"));
498    }
499
500    #[test]
501    fn record_batch_to_dataframe_zero_rows_succeeds() {
502        let schema = Arc::new(ArrowSchema::new(vec![Field::new(
503            "v",
504            DataType::Int64,
505            false,
506        )]));
507        let mut ii = Int64Builder::new();
508        let arr: ArrayRef = Arc::new(ii.finish());
509        let rb = RecordBatch::try_new(schema, vec![arr]).unwrap();
510        let df = record_batch_to_dataframe(rb).unwrap();
511        assert_eq!(df.height(), 0);
512        assert_eq!(df.width(), 1);
513    }
514
515    /// Every QuestDB table carries a designated TIMESTAMP, which the
516    /// egress decoder maps to the tz-aware Arrow `Timestamp(Microsecond,
517    /// Some("UTC"))`. Materialising that into a polars `DataFrame`
518    /// requires the `dtype-datetime` + `timezones` polars features; this
519    /// test guards that feature set (without them `Series::from_arrow`
520    /// fails at runtime, which `fetch_all_polars` on any real result set
521    /// would hit). See the `polars` feature in `Cargo.toml`.
522    #[test]
523    fn record_batch_to_dataframe_preserves_tz_timestamp() {
524        use arrow::array::TimestampMicrosecondArray;
525        let ts: ArrayRef = Arc::new(
526            TimestampMicrosecondArray::from(vec![1_700_000_000_000_000i64, 1_700_000_000_000_001])
527                .with_timezone("UTC"),
528        );
529        let schema = Arc::new(ArrowSchema::new(vec![Field::new(
530            "ts",
531            ts.data_type().clone(),
532            false,
533        )]));
534        let rb = RecordBatch::try_new(schema, vec![ts]).unwrap();
535        let df = record_batch_to_dataframe(rb).unwrap();
536        assert_eq!(df.height(), 2);
537        assert_eq!(df.width(), 1);
538        // The column must round-trip as a polars Datetime, not error out.
539        let series = df.select_at_idx(0).unwrap().as_materialized_series();
540        assert!(
541            matches!(series.dtype(), polars::prelude::DataType::Datetime(_, _)),
542            "expected polars Datetime, got {:?}",
543            series.dtype()
544        );
545    }
546
547    #[test]
548    fn record_batch_to_dataframe_symbol_dictionary_to_categorical() {
549        use arrow::array::types::UInt32Type;
550        use arrow::array::{DictionaryArray, StringArray, UInt32Array};
551
552        let values: ArrayRef = Arc::new(StringArray::from(vec!["x", "y", "z"]));
553        let keys = UInt32Array::from(vec![Some(2u32), Some(0), None, Some(1)]);
554        let dict = DictionaryArray::<UInt32Type>::new(keys, values);
555        let schema = Arc::new(ArrowSchema::new(vec![Field::new(
556            "sym",
557            DataType::Dictionary(Box::new(DataType::UInt32), Box::new(DataType::Utf8)),
558            true,
559        )]));
560        let rb = RecordBatch::try_new(schema, vec![Arc::new(dict) as ArrayRef]).unwrap();
561
562        let df = record_batch_to_dataframe(rb).unwrap();
563        assert_eq!(df.width(), 1);
564        assert_eq!(df.height(), 4);
565        let col = df.select_at_idx(0).unwrap();
566        assert_eq!(col.name().as_str(), "sym");
567        assert!(
568            matches!(col.dtype(), PlDataType::Categorical(_, _)),
569            "expected Categorical, got {:?}",
570            col.dtype()
571        );
572        let as_str = col
573            .as_materialized_series()
574            .cast(&PlDataType::String)
575            .unwrap();
576        let s = as_str.str().unwrap();
577        assert_eq!(s.get(0), Some("z"));
578        assert_eq!(s.get(1), Some("x"));
579        assert_eq!(s.get(2), None);
580        assert_eq!(s.get(3), Some("y"));
581    }
582
583    #[test]
584    fn symbol_categoricals_vstack_across_batches() {
585        // Regression guard: a SYMBOL column from one query arrives across
586        // many QWP RESULT_BATCH frames (~16k rows each), and
587        // `fetch_all_polars` vstacks the per-batch DataFrames. polars only
588        // vstacks Categoricals that share one `Categories` identity, so every
589        // batch must build against the same one. Building each batch against a
590        // fresh `Categories::random()` broke this: vstack failed with
591        // "Categories name mismatch" the moment a result spanned >1 batch.
592        use arrow::array::types::UInt32Type;
593        use arrow::array::{DictionaryArray, StringArray, UInt32Array};
594
595        fn sym_batch(values: &[&str], keys: &[Option<u32>]) -> RecordBatch {
596            let values: ArrayRef = Arc::new(StringArray::from(values.to_vec()));
597            let keys = UInt32Array::from(keys.to_vec());
598            let dict = DictionaryArray::<UInt32Type>::new(keys, values);
599            let schema = Arc::new(ArrowSchema::new(vec![Field::new(
600                "sym",
601                DataType::Dictionary(Box::new(DataType::UInt32), Box::new(DataType::Utf8)),
602                true,
603            )]));
604            RecordBatch::try_new(schema, vec![Arc::new(dict) as ArrayRef]).unwrap()
605        }
606
607        // Two batches with *different* per-batch dictionaries (the second
608        // introduces a new symbol), mirroring QWP's growing delta dictionary.
609        let b1 = sym_batch(&["a", "b"], &[Some(0), Some(1), Some(0)]); // a, b, a
610        let b2 = sym_batch(&["b", "c"], &[Some(0), Some(1)]); // b, c
611
612        let mut df = record_batch_to_dataframe(b1).unwrap();
613        let df2 = record_batch_to_dataframe(b2).unwrap();
614        df.vstack_mut_owned(df2)
615            .expect("SYMBOL Categoricals from different batches must vstack");
616
617        assert_eq!(df.height(), 5);
618        let as_str = df
619            .select_at_idx(0)
620            .unwrap()
621            .as_materialized_series()
622            .cast(&PlDataType::String)
623            .unwrap();
624        let s = as_str.str().unwrap();
625        let got: Vec<Option<&str>> = (0..df.height()).map(|i| s.get(i)).collect();
626        assert_eq!(
627            got,
628            vec![Some("a"), Some("b"), Some("a"), Some("b"), Some("c")]
629        );
630    }
631
632    #[test]
633    fn symbol_categoricals_multi_column_and_interleaved_streams() {
634        // All SYMBOL columns process-wide share `Categories::global()`, so prove
635        // the two cases that make a shared mapping suspicious stay correct:
636        // (1) several SYMBOL columns in one DataFrame, and (2) independent
637        // cursors converting interleaved. Strings are deliberately reused across
638        // columns and streams ("x", "p") to show the shared string<->id mapping
639        // never crosswires them.
640        use arrow::array::types::UInt32Type;
641        use arrow::array::{DictionaryArray, StringArray, UInt32Array};
642
643        fn sym(vals: &[&str], keys: &[Option<u32>]) -> ArrayRef {
644            let values: ArrayRef = Arc::new(StringArray::from(vals.to_vec()));
645            Arc::new(DictionaryArray::<UInt32Type>::new(
646                UInt32Array::from(keys.to_vec()),
647                values,
648            )) as ArrayRef
649        }
650        fn batch(fields: &[&str], cols: Vec<ArrayRef>) -> RecordBatch {
651            let dict_ty =
652                DataType::Dictionary(Box::new(DataType::UInt32), Box::new(DataType::Utf8));
653            let schema = Arc::new(ArrowSchema::new(
654                fields
655                    .iter()
656                    .map(|n| Field::new(*n, dict_ty.clone(), true))
657                    .collect::<Vec<_>>(),
658            ));
659            RecordBatch::try_new(schema, cols).unwrap()
660        }
661        fn vals(df: &DataFrame, i: usize) -> Vec<Option<String>> {
662            let s = df
663                .select_at_idx(i)
664                .unwrap()
665                .as_materialized_series()
666                .cast(&PlDataType::String)
667                .unwrap();
668            let s = s.str().unwrap();
669            (0..df.height())
670                .map(|r| s.get(r).map(str::to_owned))
671                .collect()
672        }
673        let some = |xs: &[&str]| xs.iter().map(|s| Some(s.to_string())).collect::<Vec<_>>();
674
675        // Stream 1: two SYMBOL columns "a","b". Stream 2: one column "a".
676        let s1b1 = batch(
677            &["a", "b"],
678            vec![
679                sym(&["x", "y"], &[Some(0), Some(1)]),
680                sym(&["m"], &[Some(0), Some(0)]),
681            ],
682        );
683        let s2b1 = batch(&["a"], vec![sym(&["x", "p"], &[Some(1), Some(0)])]);
684        let s1b2 = batch(
685            &["a", "b"],
686            vec![
687                sym(&["y", "z"], &[Some(0), Some(1)]),
688                sym(&["m", "x"], &[Some(1), Some(0)]),
689            ],
690        );
691        let s2b2 = batch(&["a"], vec![sym(&["p", "q"], &[Some(0), Some(1)])]);
692
693        // Convert interleaved (mimics two concurrent cursors), then vstack per stream.
694        let mut s1 = record_batch_to_dataframe(s1b1).unwrap();
695        let mut s2 = record_batch_to_dataframe(s2b1).unwrap();
696        s1.vstack_mut_owned(record_batch_to_dataframe(s1b2).unwrap())
697            .unwrap();
698        s2.vstack_mut_owned(record_batch_to_dataframe(s2b2).unwrap())
699            .unwrap();
700
701        assert_eq!(vals(&s1, 0), some(&["x", "y", "y", "z"])); // multi-column, col a
702        assert_eq!(vals(&s1, 1), some(&["m", "m", "x", "m"])); // col b reuses "x","m"
703        assert_eq!(vals(&s2, 0), some(&["p", "x", "p", "q"])); // other cursor, shares "x","p"
704    }
705
706    fn sym_batch(values: &[&str], keys: &[Option<u32>]) -> RecordBatch {
707        use arrow::array::types::UInt32Type;
708        use arrow::array::{DictionaryArray, StringArray, UInt32Array};
709        let values: ArrayRef = Arc::new(StringArray::from(values.to_vec()));
710        let dict = DictionaryArray::<UInt32Type>::new(UInt32Array::from(keys.to_vec()), values);
711        let schema = Arc::new(ArrowSchema::new(vec![Field::new(
712            "sym",
713            DataType::Dictionary(Box::new(DataType::UInt32), Box::new(DataType::Utf8)),
714            true,
715        )]));
716        RecordBatch::try_new(schema, vec![Arc::new(dict) as ArrayRef]).unwrap()
717    }
718
719    fn cat_strings(df: &DataFrame) -> Vec<Option<String>> {
720        let s = df
721            .select_at_idx(0)
722            .unwrap()
723            .as_materialized_series()
724            .cast(&PlDataType::String)
725            .unwrap();
726        let s = s.str().unwrap();
727        (0..df.height())
728            .map(|i| s.get(i).map(str::to_owned))
729            .collect()
730    }
731
732    #[test]
733    fn delta_symbol_registry_interns_once_and_vstacks_across_batches() {
734        // The connection dict grows; the registry registers only the new tail
735        // each batch and keys (global codes) wrap straight into a Categorical.
736        // Both batches share the registry's `Categories`, so they vstack.
737        let mut dict = SymbolDict::new();
738        dict.apply_delta(0, [b"a".as_slice(), b"b".as_slice()])
739            .unwrap();
740        let mut reg = SymbolRegistry::new();
741        reg.sync(&dict).unwrap();
742        let df1 = build_dataframe(
743            sym_batch(&["a", "b"], &[Some(0), Some(1), Some(0)]),
744            &[true],
745            &reg,
746        )
747        .unwrap();
748
749        dict.apply_delta(2, [b"c".as_slice()]).unwrap();
750        reg.sync(&dict).unwrap();
751        let df2 = build_dataframe(
752            sym_batch(&["a", "b", "c"], &[Some(2), None, Some(1)]),
753            &[true],
754            &reg,
755        )
756        .unwrap();
757
758        let mut df = df1;
759        df.vstack_mut_owned(df2)
760            .expect("registry Categoricals from different batches must vstack");
761        assert!(matches!(
762            df.select_at_idx(0).unwrap().dtype(),
763            PlDataType::Categorical(_, _)
764        ));
765        assert_eq!(
766            cat_strings(&df),
767            vec![
768                Some("a".into()),
769                Some("b".into()),
770                Some("a".into()),
771                Some("c".into()),
772                None,
773                Some("b".into()),
774            ]
775        );
776    }
777
778    #[test]
779    fn delta_symbol_registry_rebuilds_on_dict_reset() {
780        let mut dict = SymbolDict::new();
781        dict.apply_delta(0, [b"x".as_slice(), b"y".as_slice(), b"z".as_slice()])
782            .unwrap();
783        let mut reg = SymbolRegistry::new();
784        reg.sync(&dict).unwrap();
785
786        // CACHE_RESET: dict cleared then re-grown; code 0 must now be "p".
787        dict.reset();
788        dict.apply_delta(0, [b"p".as_slice()]).unwrap();
789        reg.sync(&dict).unwrap();
790
791        let df = build_dataframe(sym_batch(&["p"], &[Some(0)]), &[true], &reg).unwrap();
792        assert_eq!(cat_strings(&df), vec![Some("p".into())]);
793    }
794
795    #[test]
796    fn column_local_symbol_still_builds_via_fallback() {
797        // delta_modes = false → the column-local path (cast + take), not the
798        // registry, so a non-prefix per-batch dict is still correct.
799        let reg = SymbolRegistry::new();
800        let df = build_dataframe(
801            sym_batch(&["L0", "L1"], &[Some(1), Some(0)]),
802            &[false],
803            &reg,
804        )
805        .unwrap();
806        assert!(matches!(
807            df.select_at_idx(0).unwrap().dtype(),
808            PlDataType::Categorical(_, _)
809        ));
810        assert_eq!(cat_strings(&df), vec![Some("L1".into()), Some("L0".into())]);
811    }
812
813    #[test]
814    fn column_local_symbol_fallback_vstacks_across_batches() {
815        // The fallback interns into the registry's per-cursor `Categories`, so
816        // two column-local batches of one stream share one identity and vstack.
817        // A regression to a fresh-per-call `Categories` would fail vstack here.
818        let reg = SymbolRegistry::new();
819        let mut df =
820            build_dataframe(sym_batch(&["a", "b"], &[Some(0), Some(1)]), &[false], &reg).unwrap();
821        let df2 =
822            build_dataframe(sym_batch(&["b", "c"], &[Some(0), Some(1)]), &[false], &reg).unwrap();
823        df.vstack_mut_owned(df2)
824            .expect("fallback Categoricals from different batches must vstack");
825        assert_eq!(
826            cat_strings(&df),
827            vec![
828                Some("a".into()),
829                Some("b".into()),
830                Some("b".into()),
831                Some("c".into())
832            ]
833        );
834    }
835
836    #[test]
837    fn delta_and_fallback_use_independent_categories() {
838        // A single stream mixing a delta column (code == QWP code) and a
839        // column-local fallback column must not crosswire: the fallback's
840        // interning must not disturb the delta mapping's code alignment.
841        let mut dict = SymbolDict::new();
842        dict.apply_delta(0, [b"d0".as_slice(), b"d1".as_slice()])
843            .unwrap();
844        let mut reg = SymbolRegistry::new();
845        reg.sync(&dict).unwrap();
846
847        let dict_ty = DataType::Dictionary(Box::new(DataType::UInt32), Box::new(DataType::Utf8));
848        let schema = Arc::new(ArrowSchema::new(vec![
849            Field::new("delta", dict_ty.clone(), true),
850            Field::new("local", dict_ty, true),
851        ]));
852        use arrow::array::types::UInt32Type;
853        use arrow::array::{DictionaryArray, StringArray, UInt32Array};
854        let delta_col: ArrayRef = Arc::new(DictionaryArray::<UInt32Type>::new(
855            UInt32Array::from(vec![Some(1u32), Some(0)]),
856            Arc::new(StringArray::from(vec!["d0", "d1"])) as ArrayRef,
857        ));
858        let local_col: ArrayRef = Arc::new(DictionaryArray::<UInt32Type>::new(
859            UInt32Array::from(vec![Some(0u32), Some(1)]),
860            Arc::new(StringArray::from(vec!["d0", "L1"])) as ArrayRef,
861        ));
862        let rb = RecordBatch::try_new(schema, vec![delta_col, local_col]).unwrap();
863        let df = build_dataframe(rb, &[true, false], &reg).unwrap();
864
865        let delta = df
866            .select_at_idx(0)
867            .unwrap()
868            .as_materialized_series()
869            .cast(&PlDataType::String)
870            .unwrap();
871        let local = df
872            .select_at_idx(1)
873            .unwrap()
874            .as_materialized_series()
875            .cast(&PlDataType::String)
876            .unwrap();
877        assert_eq!(delta.str().unwrap().get(0), Some("d1"));
878        assert_eq!(delta.str().unwrap().get(1), Some("d0"));
879        assert_eq!(local.str().unwrap().get(0), Some("d0"));
880        assert_eq!(local.str().unwrap().get(1), Some("L1"));
881    }
882}