1use crate::exact::{calendar_without_out_of_range, stored_out_of_range};
19use polars::chunked_array::cast::CastOptions;
20use polars::prelude::*;
21
22pub fn can_leave_calendar(dtype: &DataType) -> bool {
25 matches!(
26 dtype,
27 DataType::Date | DataType::Datetime(TimeUnit::Milliseconds | TimeUnit::Microseconds, _)
28 )
29}
30
31pub fn text_or_stored(
34 series: &Series,
35 text: impl Fn(&Series) -> PolarsResult<StringChunked>,
36) -> PolarsResult<Series> {
37 let dtype = series.dtype();
38 let text = match calendar_without_out_of_range(series)? {
39 Some(shown) => {
41 let stored = series.to_physical_repr().cast(&DataType::Int64)?;
42 text(&shown)?
43 .iter()
44 .zip(stored.i64()?.iter())
45 .map(|(text, v)| match text {
46 Some(text) => Some(text.to_string()),
47 None => v.and_then(|v| stored_out_of_range(dtype, v)),
48 })
49 .collect::<StringChunked>()
50 }
51 None => text(series)?,
52 };
53 Ok(text.with_name(series.name().clone()).into_series())
54}
55
56pub fn cast_text(series: &Series, options: CastOptions) -> PolarsResult<Series> {
59 if !can_leave_calendar(series.dtype()) {
60 return series.cast_with_options(&DataType::String, options);
61 }
62 text_or_stored(series, |s| Ok(s.cast(&DataType::String)?.str()?.clone()))
63}
64
65pub fn formatted(series: &Series, format: &str) -> PolarsResult<Series> {
68 if !can_leave_calendar(series.dtype()) {
69 return TemporalMethods::to_string(series, format);
70 }
71 text_or_stored(series, |s| {
72 Ok(TemporalMethods::to_string(s, format)?.str()?.clone())
73 })
74}
75
76fn text_field(_: &Schema, field: &Field) -> PolarsResult<Field> {
77 Ok(Field::new(field.name().clone(), DataType::String))
78}
79
80pub fn text_expr(expr: Expr, options: CastOptions) -> Expr {
82 expr.map_with_fmt_str(
83 move |c| cast_text(c.as_materialized_series(), options).map(Column::from),
84 text_field,
85 "past_calendar_text",
86 )
87}
88
89pub fn format_expr(expr: Expr, format: String) -> Expr {
91 expr.map_with_fmt_str(
92 move |c| formatted(c.as_materialized_series(), &format).map(Column::from),
93 text_field,
94 "past_calendar_format",
95 )
96}
97
98pub fn calendar_expr(expr: Expr) -> Expr {
101 expr.map_with_fmt_str(
102 |c| {
103 Ok(calendar_without_out_of_range(c.as_materialized_series())?
104 .map(Column::from)
105 .unwrap_or(c))
106 },
107 |_, field| Ok(field.clone()),
108 "past_calendar_null",
109 )
110}
111
112const NS_PER_DAY: i128 = 86_400_000_000_000;
113
114struct Reach {
117 back: i128,
118 forward: i128,
119 low: i128,
120}
121
122fn ns_reach(function: &TemporalFunction, zoned: bool, args: &[Column]) -> Option<Reach> {
130 let spans = || -> Option<Vec<(i128, bool, bool)>> {
132 let arg = args.first()?.unique().ok()?;
134 let text = arg.as_materialized_series().str().ok()?;
135 text.iter()
136 .flatten()
137 .map(|text| {
138 let d = Duration::try_parse(text).ok()?;
139 let span = i128::from(d.months().abs()) * 31 * NS_PER_DAY
140 + i128::from(d.weeks().abs()) * 7 * NS_PER_DAY
141 + i128::from(d.days().abs()) * NS_PER_DAY
142 + i128::from(d.nanoseconds().abs());
143 Some((span, d.negative(), d.months() != 0))
144 })
145 .collect()
146 };
147 let longest = |spans: &[(i128, bool, bool)], back: Option<bool>| {
148 spans
149 .iter()
150 .filter(|(_, negative, _)| back.is_none_or(|back| *negative == back))
151 .map(|(span, _, _)| *span)
152 .max()
153 .unwrap_or(0)
154 };
155 let in_months = |spans: &[(i128, bool, bool)]| spans.iter().any(|(_, _, months)| *months);
156 let (back, forward, calendar) = match function {
157 TemporalFunction::MonthStart => (31 * NS_PER_DAY, 0, true),
158 TemporalFunction::MonthEnd => (31 * NS_PER_DAY, 32 * NS_PER_DAY, true),
160 TemporalFunction::Truncate => {
161 let spans = spans()?;
162 (longest(&spans, None), 0, in_months(&spans))
163 }
164 TemporalFunction::Round => {
166 let spans = spans()?;
167 let every = longest(&spans, None);
168 (every, every, in_months(&spans))
169 }
170 #[cfg(feature = "sql")]
172 TemporalFunction::OffsetBy => {
173 let spans = spans()?;
174 let back = longest(&spans, Some(true));
175 (back, longest(&spans, Some(false)), in_months(&spans))
176 }
177 TemporalFunction::Date
179 | TemporalFunction::Time
180 | TemporalFunction::OrdinalDay
181 | TemporalFunction::IsoYear
182 | TemporalFunction::IsLeapYear
183 | TemporalFunction::DaysInMonth
184 | TemporalFunction::Datetime
185 if zoned =>
186 {
187 (0, 0, true)
188 }
189 _ => return None,
190 };
191 let zone = if zoned { NS_PER_DAY } else { 0 };
192 Some(Reach {
193 back: back + zone,
194 forward: forward + zone,
195 low: if calendar || zoned {
196 -9_223_372_036_000_000_000
197 } else {
198 i128::from(i64::MIN)
199 },
200 })
201}
202
203fn ns_within_reach(function: &TemporalFunction, cols: &mut [Column]) -> PolarsResult<Column> {
208 let value = std::mem::take(&mut cols[0]);
209 let DataType::Datetime(TimeUnit::Nanoseconds, zone) = value.dtype() else {
210 return Ok(value);
211 };
212 let Some(reach) = ns_reach(function, zone.is_some(), &cols[1..]) else {
213 return Ok(value);
214 };
215 let fits = |v: i64| {
216 i128::from(v) - reach.back >= reach.low
217 && i128::from(v) + reach.forward <= i128::from(i64::MAX)
218 };
219 let series = value.as_materialized_series();
220 let stored = series.to_physical_repr();
221 let stored = stored.i64()?;
222 if [stored.min(), stored.max()].into_iter().flatten().all(fits) {
224 return Ok(value);
225 }
226 let kept = stored.apply(|v| v.filter(|v| fits(*v)));
227 Ok(kept
228 .into_series()
229 .cast(value.dtype())?
230 .with_name(series.name().clone())
231 .into_column())
232}
233
234fn ns_edge_expr(input: &[Expr], function: TemporalFunction) -> Expr {
238 input[0].clone().map_many(
239 move |cols| ns_within_reach(&function, cols),
240 &input[1..],
241 |_, fields| Ok(fields[0].clone()),
242 )
243}
244
245fn moves_ns(function: &TemporalFunction) -> bool {
248 #[cfg(feature = "sql")]
249 if matches!(function, TemporalFunction::OffsetBy) {
250 return true;
251 }
252 matches!(
253 function,
254 TemporalFunction::MonthStart
255 | TemporalFunction::MonthEnd
256 | TemporalFunction::Truncate
257 | TemporalFunction::Round
258 | TemporalFunction::Date
259 | TemporalFunction::Time
260 | TemporalFunction::OrdinalDay
261 | TemporalFunction::IsoYear
262 | TemporalFunction::IsLeapYear
263 | TemporalFunction::DaysInMonth
264 | TemporalFunction::Datetime
265 )
266}
267
268fn countable_dates(column: Column, unit: TimeUnit) -> PolarsResult<Column> {
274 if column.dtype() != &DataType::Date {
275 return Ok(column);
276 }
277 let per_day: i64 = match unit {
278 TimeUnit::Nanoseconds => 86_400_000_000_000,
279 TimeUnit::Microseconds => 86_400_000_000,
280 TimeUnit::Milliseconds => 86_400_000,
281 };
282 let fits = |days: i32| i64::from(days).abs() <= i64::MAX / per_day;
283 let series = column.as_materialized_series();
284 let days = series.to_physical_repr();
285 let days = days.i32()?;
286 if [days.min(), days.max()].into_iter().flatten().all(fits) {
287 return Ok(column);
288 }
289 Ok(days
290 .apply(|days| days.filter(|days| fits(*days)))
291 .into_date()
292 .into_series()
293 .with_name(series.name().clone())
294 .into_column())
295}
296
297fn countable_expr(expr: Expr, unit: TimeUnit) -> Expr {
300 expr.map_with_fmt_str(
301 move |c| countable_dates(c, unit),
302 |_, field| Ok(field.clone()),
303 "past_calendar_countable",
304 )
305}
306
307fn reads_calendar(function: &TemporalFunction) -> bool {
310 !matches!(
311 function,
312 TemporalFunction::TimeStamp(_)
313 | TemporalFunction::CastTimeUnit(_)
314 | TemporalFunction::WithTimeUnit(_)
315 | TemporalFunction::ConvertTimeZone(_)
316 )
317}
318
319pub fn guard_expr(expr: Expr, schema: Option<&Schema>) -> Expr {
333 let may_leave = |e: &Expr| match (e, schema) {
334 (Expr::Literal(_), _) => false,
335 (e, Some(schema)) => e
336 .to_field(schema)
337 .map_or(true, |f| can_leave_calendar(f.dtype())),
338 (_, None) => true,
339 };
340 let may_be_ns = |e: &Expr| match (e, schema) {
341 (Expr::Literal(_), _) => false,
342 (e, Some(schema)) => e.to_field(schema).map_or(true, |f| {
343 matches!(f.dtype(), DataType::Datetime(TimeUnit::Nanoseconds, _))
344 }),
345 (_, None) => true,
346 };
347 let as_text = |e: Expr| {
348 if may_leave(&e) {
349 text_expr(e, CastOptions::NonStrict)
350 } else {
351 e
352 }
353 };
354 let gives_text = |e: &Expr| {
355 schema.is_some_and(|schema| {
356 e.to_field(schema)
357 .is_ok_and(|f| f.dtype() == &DataType::String)
358 })
359 };
360 let may_be_date = |e: &Expr| match (e, schema) {
361 (Expr::Literal(_), _) => false,
362 (e, Some(schema)) => e
363 .to_field(schema)
364 .map_or(true, |f| f.dtype() == &DataType::Date),
365 (_, None) => true,
366 };
367 let countable = |e: Expr, unit: TimeUnit| {
368 if may_be_date(&e) {
369 countable_expr(e, unit)
370 } else {
371 e
372 }
373 };
374 let datetime_unit = |e: &Expr| match e.to_field(schema?).ok()?.dtype() {
375 DataType::Datetime(unit, _) => Some(*unit),
376 _ => None,
377 };
378 expr.map_expr(|e| match e {
379 e @ (Expr::Ternary { .. }
380 | Expr::Function {
381 function: FunctionExpr::Coalesce,
382 ..
383 }) if gives_text(&e) => match e {
384 Expr::Ternary {
385 predicate,
386 truthy,
387 falsy,
388 } => Expr::Ternary {
389 predicate,
390 truthy: Arc::new(as_text(Arc::unwrap_or_clone(truthy))),
391 falsy: Arc::new(as_text(Arc::unwrap_or_clone(falsy))),
392 },
393 Expr::Function { input, function } => Expr::Function {
394 input: input.into_iter().map(as_text).collect(),
395 function,
396 },
397 e => e,
398 },
399 e @ (Expr::Ternary { .. }
401 | Expr::Function {
402 function:
403 FunctionExpr::Coalesce
404 | FunctionExpr::FillNull
405 | FunctionExpr::MaxHorizontal
406 | FunctionExpr::MinHorizontal,
407 ..
408 }) => match (datetime_unit(&e), e) {
409 (
410 Some(unit),
411 Expr::Ternary {
412 predicate,
413 truthy,
414 falsy,
415 },
416 ) => Expr::Ternary {
417 predicate,
418 truthy: Arc::new(countable(Arc::unwrap_or_clone(truthy), unit)),
419 falsy: Arc::new(countable(Arc::unwrap_or_clone(falsy), unit)),
420 },
421 (Some(unit), Expr::Function { input, function }) => Expr::Function {
422 input: input.into_iter().map(|e| countable(e, unit)).collect(),
423 function,
424 },
425 (_, e) => e,
426 },
427 Expr::Cast {
428 expr,
429 dtype,
430 options,
431 } => match dtype.as_literal() {
432 Some(DataType::String) if may_leave(&expr) => {
433 text_expr(Arc::unwrap_or_clone(expr), options)
434 }
435 Some(DataType::Datetime(unit, _)) if may_be_date(&expr) => Expr::Cast {
436 expr: Arc::new(countable_expr(Arc::unwrap_or_clone(expr), *unit)),
437 dtype,
438 options,
439 },
440 _ => Expr::Cast {
441 expr,
442 dtype,
443 options,
444 },
445 },
446 Expr::Function {
447 mut input,
448 function: FunctionExpr::TemporalExpr(function),
449 } if input
450 .first()
451 .is_some_and(|e| may_leave(e) || (moves_ns(&function) && may_be_ns(e))) =>
452 {
453 if may_leave(&input[0]) {
454 if let TemporalFunction::ToString(format) = function {
455 return format_expr(input.swap_remove(0), format);
456 }
457 if reads_calendar(&function) {
458 input[0] = calendar_expr(input[0].clone());
459 }
460 }
461 if moves_ns(&function) && may_be_ns(&input[0]) {
462 input[0] = ns_edge_expr(&input, function.clone());
463 }
464 Expr::Function {
465 input,
466 function: FunctionExpr::TemporalExpr(function),
467 }
468 }
469 #[cfg(feature = "sql")]
471 Expr::Function {
472 input,
473 function:
474 function @ FunctionExpr::StringExpr(
475 StringFunction::ConcatHorizontal { .. } | StringFunction::ConcatVertical { .. },
476 ),
477 } => Expr::Function {
478 input: input.into_iter().map(as_text).collect(),
479 function,
480 },
481 e => e,
482 })
483}
484
485#[cfg(feature = "sql")]
487fn guards(expr: &Expr) -> bool {
488 expr.into_iter().any(|e| match e {
489 Expr::Cast { dtype, .. } => matches!(
490 dtype.as_literal(),
491 Some(DataType::String | DataType::Datetime(..))
492 ),
493 Expr::Ternary { .. } => true,
494 Expr::Function { function, .. } => matches!(
495 function,
496 FunctionExpr::TemporalExpr(_)
497 | FunctionExpr::Coalesce
498 | FunctionExpr::FillNull
499 | FunctionExpr::MaxHorizontal
500 | FunctionExpr::MinHorizontal
501 | FunctionExpr::StringExpr(
502 StringFunction::ConcatHorizontal { .. } | StringFunction::ConcatVertical { .. }
503 )
504 ),
505 _ => false,
506 })
507}
508
509#[cfg(feature = "sql")]
511fn holds_guards(plan: &DslPlan) -> bool {
512 plan.into_iter().any(|node| match node {
513 DslPlan::Filter { predicate, .. } => guards(predicate),
514 DslPlan::Select { expr, .. } => expr.iter().any(guards),
515 DslPlan::HStack { exprs, .. } => exprs.iter().any(guards),
516 DslPlan::Sort { by_column, .. } => by_column.iter().any(guards),
517 DslPlan::GroupBy {
518 keys,
519 predicates,
520 aggs,
521 ..
522 } => keys.iter().chain(predicates).chain(aggs).any(guards),
523 DslPlan::Join {
524 left_on,
525 right_on,
526 predicates,
527 ..
528 } => left_on.iter().chain(right_on).chain(predicates).any(guards),
529 DslPlan::Union { args, .. } => args.to_supertypes,
531 _ => false,
532 })
533}
534
535#[cfg(feature = "sql")]
539fn union_text(input: &mut DslPlan, union: &Schema, diagonal: bool) {
540 let Ok(own) = LazyFrame::from(input.clone()).collect_schema() else {
541 return;
542 };
543 let texts: Vec<Expr> = own
544 .iter()
545 .enumerate()
546 .filter(|(i, (name, dtype))| {
547 let stacked = if diagonal {
548 union.get(name.as_str())
549 } else {
550 union.get_at_index(*i).map(|(_, dtype)| dtype)
551 };
552 can_leave_calendar(dtype) && stacked == Some(&DataType::String)
553 })
554 .map(|(_, (name, _))| text_expr(Expr::Column(name.clone()), CastOptions::NonStrict))
555 .collect();
556 if !texts.is_empty() {
557 *input = LazyFrame::from(input.clone())
558 .with_columns(texts)
559 .logical_plan;
560 }
561}
562
563#[cfg(feature = "sql")]
567pub fn guard_plan(plan: &mut DslPlan) {
568 if !holds_guards(plan) {
569 return;
570 }
571 let schema = |input: &DslPlan| LazyFrame::from(input.clone()).collect_schema().ok();
574 let guard = |exprs: &mut Vec<Expr>, schema: Option<&Schema>| {
575 for expr in exprs.iter_mut().filter(|e| guards(e)) {
576 *expr = guard_expr(std::mem::take(expr), schema);
577 }
578 };
579 match plan {
580 DslPlan::Filter { input, predicate } if guards(predicate) => {
581 *predicate = guard_expr(std::mem::take(predicate), schema(input).as_deref());
582 }
583 DslPlan::Select { input, expr, .. } if expr.iter().any(guards) => {
584 guard(expr, schema(input).as_deref());
585 }
586 DslPlan::HStack { input, exprs, .. } if exprs.iter().any(guards) => {
587 guard(exprs, schema(input).as_deref());
588 }
589 DslPlan::Sort {
590 input, by_column, ..
591 } if by_column.iter().any(guards) => {
592 guard(by_column, schema(input).as_deref());
593 }
594 DslPlan::GroupBy {
595 input,
596 keys,
597 predicates,
598 aggs,
599 ..
600 } if keys.iter().chain(&*predicates).chain(&*aggs).any(guards) => {
601 let schema = schema(input);
602 guard(keys, schema.as_deref());
603 guard(predicates, schema.as_deref());
604 guard(aggs, schema.as_deref());
605 }
606 DslPlan::Join {
607 input_left,
608 input_right,
609 left_on,
610 right_on,
611 predicates,
612 ..
613 } => {
614 if left_on.iter().any(guards) {
615 guard(left_on, schema(input_left).as_deref());
616 }
617 if right_on.iter().any(guards) {
618 guard(right_on, schema(input_right).as_deref());
619 }
620 guard(predicates, None);
622 }
623 DslPlan::Union { args, .. } if args.to_supertypes => {
624 let union = schema(plan);
625 if let (Some(union), DslPlan::Union { inputs, args }) = (union, &mut *plan) {
626 for input in inputs {
627 union_text(input, &union, args.diagonal);
628 }
629 }
630 }
631 _ => {}
632 }
633 if let DslPlan::IR { dsl, .. } = plan {
634 let mut inner = Arc::unwrap_or_clone(dsl.clone());
637 guard_plan(&mut inner);
638 *plan = inner;
639 return;
640 }
641 crate::table::for_each_input(plan, &mut guard_plan);
642}
643
644#[cfg(test)]
645mod tests;