Skip to main content

polars_ops/series/ops/
strings.rs

1use std::borrow::Cow;
2
3use arrow::array::builder::StaticArrayBuilder;
4use arrow::array::{Array, Utf8ViewArrayBuilder};
5use arrow::datatypes::ArrowDataType;
6use polars_core::prelude::{Column, DataType, IntoColumn, StringChunked};
7use polars_core::scalar::Scalar;
8use polars_error::{PolarsContext, PolarsResult};
9use polars_utils::broadcast::broadcast_len;
10use polars_utils::pl_str::PlSmallStr;
11
12#[inline(always)]
13fn opt_str_to_string(s: Option<&str>) -> &str {
14    s.unwrap_or("null")
15}
16
17pub fn str_format(cs: &mut [Column], format: &str, insertions: &[usize]) -> PolarsResult<Column> {
18    assert_eq!(cs.len(), insertions.len());
19    assert!(!cs.is_empty()); // Checked at IR construction
20
21    let output_name = cs[0].name().clone();
22    let output_length = broadcast_len(cs.iter()).context("str.format")?;
23
24    let mut validity = None;
25    let mut num_scalar_inputs = 0;
26    for c in cs.iter_mut() {
27        if let Some(c_validity) = c.rechunk_validity() {
28            // Column with only nulls means output is only nulls.
29            if c.null_count() == c.len() {
30                return Ok(Column::full_null(
31                    output_name,
32                    output_length,
33                    &DataType::String,
34                ));
35            }
36
37            match &mut validity {
38                v @ None => *v = Some(c_validity),
39                Some(v) => *v = arrow::bitmap::and(v, &c_validity),
40            }
41        }
42
43        *c = c.cast(&DataType::String)?;
44        num_scalar_inputs += usize::from(c.len() == 1);
45    }
46
47    let mut format = Cow::Borrowed(format);
48    let mut insertions = Cow::Borrowed(insertions);
49
50    // Fill in any constants into the format string.
51    if num_scalar_inputs > 0 {
52        let mut filled_format = String::new();
53        filled_format.push_str(&format[..*insertions.first().unwrap()]);
54        insertions = Cow::Owned(
55            cs.iter()
56                .enumerate()
57                .filter_map(|(i, c)| {
58                    let v = if c.len() == 1 {
59                        filled_format.push_str(opt_str_to_string(c.str().unwrap().get(0)));
60                        None
61                    } else {
62                        Some(filled_format.len())
63                    };
64
65                    let s = if i == cs.len() - 1 {
66                        &format[insertions[i]..]
67                    } else {
68                        &format[insertions[i]..insertions[i + 1]]
69                    };
70                    filled_format.push_str(s);
71
72                    v
73                })
74                .collect(),
75        );
76        format = filled_format.into();
77    }
78
79    let format = format.as_ref();
80    let insertions = insertions.as_ref();
81
82    // If the format string is constant.
83    if num_scalar_inputs == cs.len() {
84        let sc = Scalar::from(PlSmallStr::from_str(format));
85        return Ok(Column::new_scalar(output_name, sc, output_length));
86    }
87
88    let mut builder = Utf8ViewArrayBuilder::new(ArrowDataType::Utf8View);
89    builder.reserve(output_length);
90
91    let mut arrays = cs
92        .iter()
93        .filter(|c| c.len() != 1)
94        .map(|c| {
95            let ca = c.str().unwrap();
96            let mut iter = ca.downcast_iter();
97            let arr = iter.next().unwrap();
98            (iter, arr, 0)
99        })
100        .collect::<Vec<_>>();
101
102    // @Performance. There is some smarter stuff that can be done with views and stuff. Don't think
103    // it is worth the complexity.
104
105    // Amortize the format string allocation.
106    let mut s = String::new();
107    for i in 0..output_length {
108        if validity
109            .as_ref()
110            .is_some_and(|v| !unsafe { v.get_bit_unchecked(i) })
111        {
112            unsafe { builder.push_inline_view_ignore_validity(Default::default()) };
113
114            for (iter, arr, elem_idx) in arrays.iter_mut() {
115                *elem_idx += 1;
116                if i + 1 != output_length && *elem_idx == arr.len() {
117                    *arr = iter.next().unwrap();
118                    *elem_idx = 0;
119                }
120            }
121
122            continue;
123        }
124
125        s.clear();
126        s.push_str(&format[..insertions[0]]);
127
128        for (j, (iter, arr, elem_idx)) in arrays.iter_mut().enumerate() {
129            s.push_str(opt_str_to_string(arr.get(*elem_idx)));
130            let start = insertions[j];
131            let end = insertions.get(j + 1).copied().unwrap_or(format.len());
132            s.push_str(&format[start..end]);
133
134            *elem_idx += 1;
135            if i + 1 != output_length && *elem_idx == arr.len() {
136                *arr = iter.next().unwrap();
137                *elem_idx = 0;
138            }
139        }
140
141        builder.push_value_ignore_validity(&s);
142    }
143
144    let array = builder.freeze().with_validity(validity).to_boxed();
145    Ok(unsafe { StringChunked::from_chunks(output_name, vec![array]) }.into_column())
146}