polars_utils/
broadcast.rs1use polars_error::{PolarsResult, polars_bail};
2
3pub trait BroadcastLength {
7 fn _broadcast_len(&self) -> usize;
8 fn _column_name(&self) -> Option<&str>;
9}
10
11impl<T: BroadcastLength> BroadcastLength for &T {
12 fn _broadcast_len(&self) -> usize {
13 (*self)._broadcast_len()
14 }
15
16 fn _column_name(&self) -> Option<&str> {
17 (*self)._column_name()
18 }
19}
20
21impl BroadcastLength for usize {
22 fn _broadcast_len(&self) -> usize {
23 *self
24 }
25
26 fn _column_name(&self) -> Option<&str> {
27 None
28 }
29}
30
31pub fn broadcast_len(iter: impl IntoIterator<Item = impl BroadcastLength>) -> PolarsResult<usize> {
37 let mut iter = iter.into_iter();
38 let Some(first) = iter.next() else {
39 return Ok(0);
40 };
41
42 let mut broadcast_len = first._broadcast_len();
43 let mut broadcast_val = first;
44 for val in iter {
45 let len = val._broadcast_len();
46 if broadcast_len == 1 {
47 broadcast_len = len;
48 broadcast_val = val;
49 } else if len != broadcast_len && len != 1 {
50 let fmt_opt_name = |opt_n: Option<&str>| {
51 opt_n
52 .filter(|n| !n.is_empty())
53 .map(|n| format!(" (column '{n}')"))
54 .unwrap_or_default()
55 };
56 let our_info = fmt_opt_name(val._column_name());
57 let broadcast_info = fmt_opt_name(broadcast_val._column_name());
58 polars_bail!(
59 ShapeMismatch:
60 "can't compute broadcast length, found incompatible lengths {len}{our_info} and {broadcast_len}{broadcast_info}"
61 )
62 }
63 }
64
65 Ok(broadcast_len)
66}