1use polars_arrow::legacy::time_zone::Tz;
2use polars_arrow::temporal_conversions::MILLISECONDS_IN_DAY;
3use polars_core::prelude::arity::broadcast_try_binary_elementwise;
4use polars_core::prelude::*;
5use polars_defs::time::duration::Duration;
6use polars_utils::cache::LruCache;
7
8use crate::prelude::*;
9use crate::truncate::fast_truncate;
10
11#[inline(always)]
12fn fast_round(t: i64, every: i64) -> i64 {
13 fast_truncate(t + every / 2, every)
14}
15
16pub trait PolarsRound {
17 fn round(&self, every: &StringChunked, tz: Option<&Tz>) -> PolarsResult<Self>
18 where
19 Self: Sized;
20}
21
22impl PolarsRound for DatetimeChunked {
23 fn round(&self, every: &StringChunked, tz: Option<&Tz>) -> PolarsResult<Self> {
24 let time_zone = self.time_zone();
25 let offset = Duration::new(0);
26
27 if every.len() == 1 {
29 if let Some(every) = every.get(0) {
30 let every_parsed = Duration::try_parse(every)?;
31 if every_parsed.negative {
32 polars_bail!(ComputeError: "cannot round a Datetime to a negative duration")
33 }
34 if (time_zone.is_none() || time_zone == &Some(TimeZone::UTC))
35 && (every_parsed.months() == 0 && every_parsed.weeks() == 0)
36 {
37 let every = every_parsed.duration(self.time_unit());
40 return Ok(self
41 .physical()
42 .apply_values(|t| fast_round(t, every))
43 .into_datetime(self.time_unit(), time_zone.clone()));
44 } else {
45 let w = Window::new(every_parsed, every_parsed, offset);
46 let out = self
47 .physical()
48 .try_apply_nonnull_values_generic(|t| w.round(self.time_unit(), t, tz));
49 return Ok(out?.into_datetime(self.time_unit(), self.time_zone().clone()));
50 }
51 } else {
52 return Ok(Int64Chunked::full_null(self.name().clone(), self.len())
53 .into_datetime(self.time_unit(), self.time_zone().clone()));
54 }
55 }
56
57 polars_ensure!(
58 self.len() == every.len() || self.len() == 1,
59 length_mismatch = "dt.round",
60 self.len(),
61 every.len()
62 );
63
64 let mut duration_cache = LruCache::with_capacity((every.len() as f64).sqrt() as usize);
66
67 let out = broadcast_try_binary_elementwise(
68 self.physical(),
69 every,
70 |opt_timestamp, opt_every| match (opt_timestamp, opt_every) {
71 (Some(timestamp), Some(every)) => {
72 let every = *duration_cache.get_or_insert_with(every, Duration::parse);
73
74 if every.negative {
75 polars_bail!(ComputeError: "cannot round a Datetime to a negative duration")
76 }
77
78 let w = Window::new(every, every, offset);
79 w.round(self.time_unit(), timestamp, tz).map(Some)
80 },
81 _ => Ok(None),
82 },
83 );
84 Ok(out?.into_datetime(self.time_unit(), self.time_zone().clone()))
85 }
86}
87
88impl PolarsRound for DateChunked {
89 fn round(&self, every: &StringChunked, _tz: Option<&Tz>) -> PolarsResult<Self> {
90 let offset = Duration::new(0);
91 let out = match every.len() {
92 1 => {
93 if let Some(every) = every.get(0) {
94 let every = Duration::try_parse(every)?;
95 if every.negative {
96 polars_bail!(ComputeError: "cannot round a Date to a negative duration")
97 }
98 let w = Window::new(every, every, offset);
99 self.physical().try_apply_nonnull_values_generic(|t| {
100 Ok((w.round(
101 TimeUnit::Milliseconds,
102 MILLISECONDS_IN_DAY * t as i64,
103 None,
104 )? / MILLISECONDS_IN_DAY) as i32)
105 })
106 } else {
107 Ok(Int32Chunked::full_null(self.name().clone(), self.len()))
108 }
109 },
110 _ => {
111 polars_ensure!(
112 self.len() == every.len() || self.len() == 1,
113 length_mismatch = "dt.round",
114 self.len(),
115 every.len()
116 );
117 let mut duration_cache =
119 LruCache::with_capacity((every.len() as f64).sqrt() as usize);
120 broadcast_try_binary_elementwise(self.physical(), every, |opt_t, opt_every| match (
121 opt_t, opt_every,
122 ) {
123 (Some(t), Some(every)) => {
124 let every = *duration_cache.get_or_insert_with(every, Duration::parse);
125
126 if every.negative {
127 polars_bail!(ComputeError: "cannot round a Date to a negative duration")
128 }
129
130 let w = Window::new(every, every, offset);
131 Ok(Some(
132 (w.round(TimeUnit::Milliseconds, MILLISECONDS_IN_DAY * t as i64, None)?
133 / MILLISECONDS_IN_DAY) as i32,
134 ))
135 },
136 _ => Ok(None),
137 })
138 },
139 };
140 Ok(out?.into_date())
141 }
142}