Skip to main content

polars_utils/
broadcast.rs

1use polars_error::{PolarsResult, polars_bail};
2
3// Never intended to be in-scope or used, but allows us to be generic over types
4// we can put into broadcast_len, and provide better error messages if they're
5// not just usize lengths.
6pub 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
31// Calculates the common length a set of lengths should be broadcast to. This
32// is the shared non-unit length of the set. If there are multiple different
33// non-unit lengths this returns an error.
34//
35// Returns 0 if the iterator is empty.
36pub 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}