blob: e56816aed308b5db05834c744666d49a8a20c356 [file]
// Licensed to the Apache Software Foundation (ASF) under one
// or more contributor license agreements. See the NOTICE file
// distributed with this work for additional information
// regarding copyright ownership. The ASF licenses this file
// to you under the Apache License, Version 2.0 (the
// "License"); you may not use this file except in compliance
// with the License. You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing,
// software distributed under the License is distributed on an
// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
// KIND, either express or implied. See the License for the
// specific language governing permissions and limitations
// under the License.
use crate::conversion_funcs::trim::{trim_all, trim_all_bytes, trim_all_range, trim_java_string};
use crate::{EvalMode, SparkError, SparkResult};
use arrow::array::timezone::Tz;
use arrow::array::{
Array, ArrayRef, ArrowPrimitiveType, BooleanArray, Decimal128Builder, GenericStringArray,
OffsetSizeTrait, PrimitiveArray, PrimitiveBuilder, StringArray,
};
use arrow::datatypes::{
i256, is_validate_decimal_precision, DataType, Date32Type, Decimal256Type, Float32Type,
Float64Type, Int16Type, Int32Type, Int64Type, Int8Type, TimestampMicrosecondType,
};
use chrono::{LocalResult, NaiveDate, NaiveTime, Offset, TimeZone, Timelike};
use num::traits::CheckedNeg;
use num::{CheckedSub, Integer};
use regex::Regex;
use std::num::Wrapping;
use std::str::FromStr;
use std::sync::{Arc, LazyLock};
// Shared macro for casting UTF-8 string arrays to timestamp types (both TZ and NTZ).
// $builder is a PrimitiveBuilder expression; $extra_args are forwarded to $cast_method
// after (value, eval_mode).
macro_rules! cast_utf8_to_timestamp {
($array:expr, $eval_mode:expr, $builder:expr, $cast_method:ident $(, $extra_arg:expr)*) => {{
let mut cast_array = $builder;
let mut cast_err: Option<SparkError> = None;
for i in 0..$array.len() {
if $array.is_null(i) {
cast_array.append_null()
} else {
// trim_end only: leading spaces affect parsing (e.g. " T2" -> null, "T2" -> valid)
match $cast_method($array.value(i).trim_end(), $eval_mode $(, $extra_arg)*) {
Ok(Some(cast_value)) => cast_array.append_value(cast_value),
Ok(None) => cast_array.append_null(),
Err(e) => {
if $eval_mode == EvalMode::Ansi {
let raw_value = $array.value(i).to_string();
let e = match e {
SparkError::InvalidInputInCastToDatetime {
from_type,
to_type,
..
} => SparkError::InvalidInputInCastToDatetime {
value: raw_value,
from_type,
to_type,
},
other => other,
};
cast_err = Some(e);
break;
}
cast_array.append_null()
}
}
}
}
if let Some(e) = cast_err {
Err(e)
} else {
Ok(Arc::new(cast_array.finish()) as ArrayRef)
}
}};
}
macro_rules! cast_utf8_to_int {
($array:expr, $array_type:ty, $parse_fn:expr) => {{
let len = $array.len();
let mut cast_array = PrimitiveArray::<$array_type>::builder(len);
let parse_fn = $parse_fn;
if $array.null_count() == 0 {
for i in 0..len {
if let Some(cast_value) = parse_fn($array.value(i))? {
cast_array.append_value(cast_value);
} else {
cast_array.append_null()
}
}
} else {
for i in 0..len {
if $array.is_null(i) {
cast_array.append_null()
} else if let Some(cast_value) = parse_fn($array.value(i))? {
cast_array.append_value(cast_value);
} else {
cast_array.append_null()
}
}
}
let result: SparkResult<ArrayRef> = Ok(Arc::new(cast_array.finish()) as ArrayRef);
result
}};
}
#[derive(Clone)]
struct TimeStampInfo {
year: i32,
month: u32,
day: u32,
hour: u32,
minute: u32,
second: u32,
microsecond: u32,
}
impl Default for TimeStampInfo {
fn default() -> Self {
TimeStampInfo {
year: 1,
month: 1,
day: 1,
hour: 0,
minute: 0,
second: 0,
microsecond: 0,
}
}
}
impl TimeStampInfo {
fn with_year(&mut self, year: i32) -> &mut Self {
self.year = year;
self
}
fn with_month(&mut self, month: u32) -> &mut Self {
self.month = month;
self
}
fn with_day(&mut self, day: u32) -> &mut Self {
self.day = day;
self
}
fn with_hour(&mut self, hour: u32) -> &mut Self {
self.hour = hour;
self
}
fn with_minute(&mut self, minute: u32) -> &mut Self {
self.minute = minute;
self
}
fn with_second(&mut self, second: u32) -> &mut Self {
self.second = second;
self
}
fn with_microsecond(&mut self, microsecond: u32) -> &mut Self {
self.microsecond = microsecond;
self
}
}
pub(crate) fn is_df_cast_from_string_spark_compatible(to_type: &DataType) -> bool {
matches!(to_type, DataType::Binary)
}
pub(crate) fn cast_string_to_float(
array: &ArrayRef,
to_type: &DataType,
eval_mode: EvalMode,
) -> SparkResult<ArrayRef> {
match to_type {
DataType::Float32 => cast_string_to_float_impl::<Float32Type>(array, eval_mode, "FLOAT"),
DataType::Float64 => cast_string_to_float_impl::<Float64Type>(array, eval_mode, "DOUBLE"),
_ => Err(SparkError::Internal(format!(
"Unsupported cast to float type: {:?}",
to_type
))),
}
}
fn cast_string_to_float_impl<T: ArrowPrimitiveType>(
array: &ArrayRef,
eval_mode: EvalMode,
type_name: &str,
) -> SparkResult<ArrayRef>
where
T::Native: FromStr + num::Float,
{
let arr = array
.as_any()
.downcast_ref::<StringArray>()
.ok_or_else(|| SparkError::Internal("Expected string array".to_string()))?;
let mut builder = PrimitiveBuilder::<T>::with_capacity(arr.len());
for i in 0..arr.len() {
if arr.is_null(i) {
builder.append_null();
} else {
// `Double.parseDouble` calls `String.trim` before parsing, so only bytes <= 0x20
// are trimmed here -- `0x7F` is not whitespace to this cast, and no non-ASCII
// whitespace is trimmed by any Spark cast.
let str_value = trim_java_string(arr.value(i));
match parse_string_to_float(str_value) {
Some(v) => builder.append_value(v),
None => {
if eval_mode == EvalMode::Ansi {
return Err(invalid_value(arr.value(i), "STRING", type_name));
}
builder.append_null();
}
}
}
}
Ok(Arc::new(builder.finish()))
}
/// helper to parse floats from string inputs
fn parse_string_to_float<F>(s: &str) -> Option<F>
where
F: FromStr + num::Float,
{
// Handle +inf / -inf
if s.eq_ignore_ascii_case("inf")
|| s.eq_ignore_ascii_case("+inf")
|| s.eq_ignore_ascii_case("infinity")
|| s.eq_ignore_ascii_case("+infinity")
{
return Some(F::infinity());
}
if s.eq_ignore_ascii_case("-inf") || s.eq_ignore_ascii_case("-infinity") {
return Some(F::neg_infinity());
}
if s.eq_ignore_ascii_case("nan") {
return Some(F::nan());
}
// Remove D/F suffix if present
let pruned_float_str =
if s.ends_with("d") || s.ends_with("D") || s.ends_with('f') || s.ends_with('F') {
&s[..s.len() - 1]
} else {
s
};
// Rust's parse logic already handles scientific notations so we just rely on it
pruned_float_str.parse::<F>().ok()
}
pub(crate) fn spark_cast_utf8_to_boolean<OffsetSize>(
from: &dyn Array,
eval_mode: EvalMode,
) -> SparkResult<ArrayRef>
where
OffsetSize: OffsetSizeTrait,
{
let array = from
.as_any()
.downcast_ref::<GenericStringArray<OffsetSize>>()
.unwrap();
let output_array = array
.iter()
.map(|value| match value {
Some(value) => match trim_all(value) {
v if is_true_string(v) => Ok(Some(true)),
v if is_false_string(v) => Ok(Some(false)),
_ if eval_mode == EvalMode::Ansi => Err(SparkError::CastInvalidValue {
value: value.to_string(),
from_type: "STRING".to_string(),
to_type: "BOOLEAN".to_string(),
}),
_ => Ok(None),
},
_ => Ok(None),
})
.collect::<Result<BooleanArray, _>>()?;
Ok(Arc::new(output_array))
}
/// Equivalent to `org.apache.spark.sql.catalyst.util.StringUtils.isTrueString`, minus the trim
/// that the caller has already applied.
///
/// Spark lowercases with `UTF8String.toLowerCase` before comparing, but every candidate is
/// ASCII, and no non-ASCII character lowercases into an ASCII one that would complete any of
/// them, so an ASCII-insensitive comparison gives the same answer without allocating.
#[inline]
fn is_true_string(trimmed: &str) -> bool {
["t", "true", "y", "yes", "1"]
.iter()
.any(|v| trimmed.eq_ignore_ascii_case(v))
}
/// Equivalent to `org.apache.spark.sql.catalyst.util.StringUtils.isFalseString`; see
/// [`is_true_string`] for why the comparison is ASCII-only.
#[inline]
fn is_false_string(trimmed: &str) -> bool {
["f", "false", "n", "no", "0"]
.iter()
.any(|v| trimmed.eq_ignore_ascii_case(v))
}
pub(crate) fn cast_string_to_decimal(
array: &ArrayRef,
to_type: &DataType,
precision: &u8,
scale: &i8,
eval_mode: EvalMode,
) -> SparkResult<ArrayRef> {
match to_type {
DataType::Decimal128(_, _) => {
cast_string_to_decimal128_impl(array, eval_mode, *precision, *scale)
}
DataType::Decimal256(_, _) => {
cast_string_to_decimal256_impl(array, eval_mode, *precision, *scale)
}
_ => Err(SparkError::Internal(format!(
"Unexpected type in cast_string_to_decimal: {:?}",
to_type
))),
}
}
fn cast_string_to_decimal128_impl(
array: &ArrayRef,
eval_mode: EvalMode,
precision: u8,
scale: i8,
) -> SparkResult<ArrayRef> {
let string_array = array
.as_any()
.downcast_ref::<StringArray>()
.ok_or_else(|| SparkError::Internal("Expected string array".to_string()))?;
let mut decimal_builder = Decimal128Builder::with_capacity(string_array.len());
for i in 0..string_array.len() {
if string_array.is_null(i) {
decimal_builder.append_null();
} else {
let str_value = string_array.value(i);
match parse_string_to_decimal(str_value, precision, scale) {
Ok(Some(decimal_value)) => {
decimal_builder.append_value(decimal_value);
}
Ok(None) => {
if eval_mode == EvalMode::Ansi {
return Err(invalid_value(
string_array.value(i),
"STRING",
&format!("DECIMAL({},{})", precision, scale),
));
}
decimal_builder.append_null();
}
Err(e) => {
if eval_mode == EvalMode::Ansi {
return Err(e);
}
decimal_builder.append_null();
}
}
}
}
Ok(Arc::new(
decimal_builder
.with_precision_and_scale(precision, scale)
.map_err(|e| {
if matches!(e, arrow::error::ArrowError::InvalidArgumentError(_))
&& e.to_string().contains("too large to store in a Decimal128")
{
// Fallback error handling
SparkError::NumericValueOutOfRange {
value: "overflow".to_string(),
precision,
scale,
}
} else {
SparkError::Arrow(Arc::new(e))
}
})?
.finish(),
))
}
fn cast_string_to_decimal256_impl(
array: &ArrayRef,
eval_mode: EvalMode,
precision: u8,
scale: i8,
) -> SparkResult<ArrayRef> {
let string_array = array
.as_any()
.downcast_ref::<StringArray>()
.ok_or_else(|| SparkError::Internal("Expected string array".to_string()))?;
let mut decimal_builder = PrimitiveBuilder::<Decimal256Type>::with_capacity(string_array.len());
for i in 0..string_array.len() {
if string_array.is_null(i) {
decimal_builder.append_null();
} else {
let str_value = string_array.value(i);
match parse_string_to_decimal(str_value, precision, scale) {
Ok(Some(decimal_value)) => {
// Convert i128 to i256
let i256_value = i256::from_i128(decimal_value);
decimal_builder.append_value(i256_value);
}
Ok(None) => {
if eval_mode == EvalMode::Ansi {
return Err(invalid_value(
str_value,
"STRING",
&format!("DECIMAL({},{})", precision, scale),
));
}
decimal_builder.append_null();
}
Err(e) => {
if eval_mode == EvalMode::Ansi {
return Err(e);
}
decimal_builder.append_null();
}
}
}
}
Ok(Arc::new(
decimal_builder
.with_precision_and_scale(precision, scale)
.map_err(|e| {
if matches!(e, arrow::error::ArrowError::InvalidArgumentError(_))
&& e.to_string().contains("too large to store in a Decimal128")
{
// Fallback error handling
SparkError::NumericValueOutOfRange {
value: "overflow".to_string(),
precision,
scale,
}
} else {
SparkError::Arrow(Arc::new(e))
}
})?
.finish(),
))
}
/// Normalize fullwidth Unicode digits (U+FF10–U+FF19) to their ASCII equivalents.
///
/// Spark's UTF8String parser treats fullwidth digits as numerically equivalent to
/// ASCII digits, e.g. "123.45" parses as 123.45. Each fullwidth digit encodes
/// to exactly three UTF-8 bytes: [0xEF, 0xBC, 0x90+n] for digit n. The ASCII
/// equivalent is 0x30+n, so the conversion is: third_byte - 0x60.
///
/// All other bytes (ASCII or other multi-byte sequences) are passed through
/// unchanged, so the output is valid UTF-8 whenever the input is.
fn normalize_fullwidth_digits(s: &str) -> String {
let bytes = s.as_bytes();
let mut out = Vec::with_capacity(s.len());
let mut i = 0;
while i < bytes.len() {
if i + 2 < bytes.len()
&& bytes[i] == 0xEF
&& bytes[i + 1] == 0xBC
&& bytes[i + 2] >= 0x90
&& bytes[i + 2] <= 0x99
{
// e.g. 0x91 - 0x60 = 0x31 = b'1'
out.push(bytes[i + 2] - 0x60);
i += 3;
} else {
out.push(bytes[i]);
i += 1;
}
}
// SAFETY: we only replace valid 3-byte UTF-8 sequences [EF BC 9X] with a
// single ASCII byte; all other bytes are copied unchanged, preserving the
// UTF-8 invariant of the input.
unsafe { String::from_utf8_unchecked(out) }
}
/// Powers of ten that fit in an `i128` (`10^0` through `10^38`).
const POW10_I128: [i128; 39] = {
let mut table = [1i128; 39];
let mut i = 1;
while i < 39 {
table[i] = table[i - 1] * 10;
i += 1;
}
table
};
/// `10^exp`, or `None` when the exponent overflows an `i128` (exp >= 39).
#[inline]
pub(crate) fn pow10_i128(exp: u32) -> Option<i128> {
POW10_I128.get(exp as usize).copied()
}
/// Divide by a power of ten with HALF_UP rounding, matching `BigDecimal.setScale`: a tie
/// rounds away from zero.
///
/// `divisor` must be `10^n` for `n >= 1`, so `divisor / 2` is exact and a zero remainder
/// can never be mistaken for a tie.
#[inline]
pub(crate) fn div_round_half_up_i128(numerator: i128, divisor: i128) -> i128 {
debug_assert!(divisor >= 10);
let quotient = numerator / divisor;
let remainder = numerator % divisor;
if remainder.abs() >= divisor / 2 {
quotient + numerator.signum()
} else {
quotient
}
}
/// Accumulate an ASCII-digit slice into an `i128`, returning `None` on overflow.
///
/// The first 38 digits always fit (`i128::MAX` is ~1.7e38), so only the digits past
/// them need the per-digit overflow checks.
#[inline]
pub(crate) fn digits_to_i128(digits: &[u8]) -> Option<i128> {
let (head, tail) = digits.split_at(digits.len().min(38));
let mut value: i128 = 0;
for &d in head {
value = value * 10 + (d - b'0') as i128;
}
for &d in tail {
value = value.checked_mul(10)?.checked_add((d - b'0') as i128)?;
}
Some(value)
}
/// Values that Spark parses as NULL rather than as a decimal, matched case-insensitively.
const SPECIAL_DECIMAL_VALUES: [&str; 7] = [
"inf",
"+inf",
"-inf",
"infinity",
"+infinity",
"-infinity",
"nan",
];
/// True if `trimmed` is one of [`SPECIAL_DECIMAL_VALUES`].
#[inline]
fn is_special_value(trimmed: &str) -> bool {
// Every special value starts with `i`/`n`, or with a sign followed by `i`, so
// ordinary numeric input is ruled out after inspecting a single byte.
let bytes = trimmed.as_bytes();
let plausible = match bytes.first() {
Some(b'i' | b'I' | b'n' | b'N') => true,
Some(b'+' | b'-') => matches!(bytes.get(1), Some(b'i' | b'I')),
_ => false,
};
plausible
&& SPECIAL_DECIMAL_VALUES
.iter()
.any(|v| trimmed.eq_ignore_ascii_case(v))
}
/// Parse a decimal string into mantissa and scale
/// e.g., "123.45" -> (12345, 2), "-0.001" -> (-1, 3) , 0e50 -> (0,50) etc
/// Parse a string to decimal following Spark's behavior
fn parse_string_to_decimal(input_str: &str, precision: u8, scale: i8) -> SparkResult<Option<i128>> {
// Spark parses via `new java.math.BigDecimal(str.toString.trim)`, so the trim set is
// `String.trim`'s: every byte <= 0x20, which includes the null byte ("123\u0000" and
// "\u0000123" both parse as 123) but not 0x7F or any non-ASCII whitespace. Null bytes in
// the middle are not trimmed and will fail the digit validation in parse_decimal_str,
// producing NULL.
let trimmed = trim_java_string(input_str);
// Normalize fullwidth digits to ASCII. Fast path skips the allocation for
// pure-ASCII strings, which is the common case.
let normalized;
let trimmed = if trimmed.is_ascii() {
trimmed
} else {
normalized = normalize_fullwidth_digits(trimmed);
normalized.as_str()
};
if trimmed.is_empty() {
return Ok(None);
}
// Handle special values (inf, nan, etc.)
if is_special_value(trimmed) {
return Ok(None);
}
// validate and parse mantissa and exponent or bubble up the error
let (mantissa, exponent) = parse_decimal_str(trimmed, input_str, precision, scale)?;
// Early return mantissa 0, Spark checks if it fits digits and throw error in ansi
if mantissa == 0 {
if exponent < -37 {
return Err(SparkError::NumericOutOfRange {
value: input_str.to_string(),
});
}
return Ok(Some(0));
}
// scale adjustment
let target_scale = scale as i32;
let scale_adjustment = target_scale - exponent;
let scaled_value = if scale_adjustment >= 0 {
// Need to multiply (increase scale) but return None if scale is too high to fit i128
if scale_adjustment > 38 {
return Ok(None);
}
// Bounded above, so pow10_i128 always returns Some.
mantissa.checked_mul(pow10_i128(scale_adjustment as u32).unwrap())
} else {
// Need to divide (decrease scale)
let abs_scale_adjustment = (-scale_adjustment) as u32;
if abs_scale_adjustment > 38 {
return Ok(Some(0));
}
// Bounded above, so pow10_i128 always returns Some. The adjustment is at least 1
// here, so the divisor is a power of ten no smaller than 10.
let divisor = pow10_i128(abs_scale_adjustment).unwrap();
Some(div_round_half_up_i128(mantissa, divisor))
};
match scaled_value {
Some(value) => {
if is_validate_decimal_precision(value, precision) {
Ok(Some(value))
} else {
// Value ok but exceeds precision mentioned . THrow error
Err(SparkError::NumericValueOutOfRange {
value: trimmed.to_string(),
precision,
scale,
})
}
}
None => {
// Overflow when scaling raise exception
Err(SparkError::NumericValueOutOfRange {
value: trimmed.to_string(),
precision,
scale,
})
}
}
}
fn invalid_decimal_cast(value: &str, precision: u8, scale: i8) -> SparkError {
invalid_value(
value,
"STRING",
&format!("DECIMAL({},{})", precision, scale),
)
}
/// Parse a decimal string into mantissa and scale
/// e.g., "123.45" -> (12345, 2), "-0.001" -> (-1, 3) , 0e50 -> (0,50) etc
fn parse_decimal_str(
s: &str,
original_str: &str,
precision: u8,
scale: i8,
) -> SparkResult<(i128, i32)> {
let bytes = s.as_bytes();
let mut pos = 0;
let negative = match bytes.first() {
Some(b'-') => {
pos = 1;
true
}
Some(b'+') => {
pos = 1;
false
}
_ => false,
};
// Single validating pass over the mantissa: ASCII digits with at most one `.`,
// ending at the optional exponent marker. It also locates the integral/fractional
// split and the start of the exponent. `.`, `e` and `E` are ASCII, so scanning
// bytes can never land inside a multi-byte character and every index taken here is
// on a char boundary.
let digits_start = pos;
let mut dot_pos = None;
let mut exp_pos = None;
while pos < bytes.len() {
match bytes[pos] {
b'0'..=b'9' => pos += 1,
b'.' if dot_pos.is_none() => {
dot_pos = Some(pos);
pos += 1;
}
b'e' | b'E' => {
exp_pos = Some(pos);
break;
}
_ => return Err(invalid_decimal_cast(original_str, precision, scale)),
}
}
let exponent: i32 = match exp_pos {
Some(e_pos) => s[e_pos + 1..]
.parse()
.map_err(|_| invalid_decimal_cast(original_str, precision, scale))?,
None => 0,
};
// An empty integral part is valid (e.g. ".5" or "-.7e9"), as is an empty
// fractional part, but they cannot both be empty.
let mantissa_end = exp_pos.unwrap_or(pos);
let (integral_part, fractional_part): (&[u8], &[u8]) = match dot_pos {
Some(dot) => (&bytes[digits_start..dot], &bytes[dot + 1..mantissa_end]),
None => (&bytes[digits_start..mantissa_end], &[]),
};
if integral_part.is_empty() && fractional_part.is_empty() {
return Err(invalid_decimal_cast(original_str, precision, scale));
}
let integral_value = digits_to_i128(integral_part)
.ok_or_else(|| invalid_decimal_cast(original_str, precision, scale))?;
let fractional_scale = fractional_part.len() as i32;
let fractional_value = digits_to_i128(fractional_part)
.ok_or_else(|| invalid_decimal_cast(original_str, precision, scale))?;
// Combine: value = integral * 10^fractional_scale + fractional.
// A fractional_scale beyond 38 cannot fit in an i128, so pow10_i128 returns None and this
// maps to the invalid-decimal error path instead of panicking.
let mantissa = pow10_i128(fractional_scale as u32)
.and_then(|p| integral_value.checked_mul(p))
.and_then(|v| v.checked_add(fractional_value))
.ok_or_else(|| invalid_decimal_cast(original_str, precision, scale))?;
let final_mantissa = if negative { -mantissa } else { mantissa };
// final scale = fractional_scale - exponent
// For example : "1.23E-5" has fractional_scale=2, exponent=-5, so scale = 2 - (-5) = 7
let final_scale = fractional_scale - exponent;
Ok((final_mantissa, final_scale))
}
pub(crate) fn cast_string_to_date(
array: &ArrayRef,
to_type: &DataType,
eval_mode: EvalMode,
) -> SparkResult<ArrayRef> {
let string_array = array
.as_any()
.downcast_ref::<GenericStringArray<i32>>()
.expect("Expected a string array");
if to_type != &DataType::Date32 {
unreachable!("Invalid data type {:?} in cast from string", to_type);
}
let len = string_array.len();
let mut cast_array = PrimitiveArray::<Date32Type>::builder(len);
for i in 0..len {
let value = if string_array.is_null(i) {
None
} else {
match date_parser(string_array.value(i), eval_mode) {
Ok(Some(cast_value)) => Some(cast_value),
Ok(None) => None,
Err(e) => return Err(e),
}
};
match value {
Some(cast_value) => cast_array.append_value(cast_value),
None => cast_array.append_null(),
}
}
Ok(Arc::new(cast_array.finish()) as ArrayRef)
}
pub(crate) fn cast_string_to_timestamp(
array: &ArrayRef,
to_type: &DataType,
eval_mode: EvalMode,
timezone_str: &str,
is_spark4_plus: bool,
) -> SparkResult<ArrayRef> {
let string_array = array
.as_any()
.downcast_ref::<GenericStringArray<i32>>()
.expect("Expected a string array");
let tz = &Tz::from_str(timezone_str)
.map_err(|_| SparkError::Internal(format!("Invalid timezone string: {timezone_str}")))?;
let cast_array: ArrayRef = match to_type {
DataType::Timestamp(_, tz_opt) => {
let to_tz = tz_opt.as_deref().unwrap_or("UTC");
cast_utf8_to_timestamp!(
string_array,
eval_mode,
PrimitiveArray::<TimestampMicrosecondType>::builder(string_array.len())
.with_timezone(to_tz),
timestamp_parser,
tz,
is_spark4_plus
)?
}
_ => unreachable!("Invalid data type {:?} in cast from string", to_type),
};
Ok(cast_array)
}
pub(crate) fn cast_string_to_timestamp_ntz(
array: &ArrayRef,
eval_mode: EvalMode,
allow_time_zone: bool,
is_spark4_plus: bool,
) -> SparkResult<ArrayRef> {
let string_array = array
.as_any()
.downcast_ref::<GenericStringArray<i32>>()
.expect("Expected a string array");
let cast_array: ArrayRef = cast_utf8_to_timestamp!(
string_array,
eval_mode,
PrimitiveArray::<TimestampMicrosecondType>::builder(string_array.len()),
timestamp_ntz_parser,
allow_time_zone,
is_spark4_plus
)?;
Ok(cast_array)
}
pub(crate) fn cast_string_to_int<OffsetSize: OffsetSizeTrait>(
to_type: &DataType,
array: &ArrayRef,
eval_mode: EvalMode,
) -> SparkResult<ArrayRef> {
let string_array = array
.as_any()
.downcast_ref::<GenericStringArray<OffsetSize>>()
.expect("cast_string_to_int expected a string array");
// Select parse function once per batch based on eval_mode
let cast_array: ArrayRef =
match (to_type, eval_mode) {
(DataType::Int8, EvalMode::Legacy) => {
cast_utf8_to_int!(string_array, Int8Type, parse_string_to_i8_legacy)?
}
(DataType::Int8, EvalMode::Ansi) => {
cast_utf8_to_int!(string_array, Int8Type, parse_string_to_i8_ansi)?
}
(DataType::Int8, EvalMode::Try) => {
cast_utf8_to_int!(string_array, Int8Type, parse_string_to_i8_try)?
}
(DataType::Int16, EvalMode::Legacy) => {
cast_utf8_to_int!(string_array, Int16Type, parse_string_to_i16_legacy)?
}
(DataType::Int16, EvalMode::Ansi) => {
cast_utf8_to_int!(string_array, Int16Type, parse_string_to_i16_ansi)?
}
(DataType::Int16, EvalMode::Try) => {
cast_utf8_to_int!(string_array, Int16Type, parse_string_to_i16_try)?
}
(DataType::Int32, EvalMode::Legacy) => cast_utf8_to_int!(
string_array,
Int32Type,
|s| do_parse_string_to_int_legacy::<i32>(s, i32::MIN)
)?,
(DataType::Int32, EvalMode::Ansi) => {
cast_utf8_to_int!(string_array, Int32Type, |s| do_parse_string_to_int_ansi::<
i32,
>(
s, "INT", i32::MIN
))?
}
(DataType::Int32, EvalMode::Try) => {
cast_utf8_to_int!(
string_array,
Int32Type,
|s| do_parse_string_to_int_try::<i32>(s, i32::MIN)
)?
}
(DataType::Int64, EvalMode::Legacy) => cast_utf8_to_int!(
string_array,
Int64Type,
|s| do_parse_string_to_int_legacy::<i64>(s, i64::MIN)
)?,
(DataType::Int64, EvalMode::Ansi) => {
cast_utf8_to_int!(string_array, Int64Type, |s| do_parse_string_to_int_ansi::<
i64,
>(
s, "BIGINT", i64::MIN
))?
}
(DataType::Int64, EvalMode::Try) => {
cast_utf8_to_int!(
string_array,
Int64Type,
|s| do_parse_string_to_int_try::<i64>(s, i64::MIN)
)?
}
(dt, _) => unreachable!(
"{}",
format!("invalid integer type {dt} in cast from string")
),
};
Ok(cast_array)
}
/// Finalizes the result by applying the sign. Returns None if overflow would occur.
fn finalize_int_result<T: Integer + CheckedNeg + Copy>(result: T, negative: bool) -> Option<T> {
if negative {
Some(result)
} else {
result.checked_neg().filter(|&n| n >= T::zero())
}
}
/// Equivalent to
/// - org.apache.spark.unsafe.types.UTF8String.toInt(IntWrapper intWrapper, boolean allowDecimal)
/// - org.apache.spark.unsafe.types.UTF8String.toLong(LongWrapper longWrapper, boolean allowDecimal)
fn do_parse_string_to_int_legacy<T: Integer + CheckedSub + CheckedNeg + From<u8> + Copy>(
str: &str,
min_value: T,
) -> SparkResult<Option<T>> {
let trimmed_bytes = trim_all_bytes(str.as_bytes());
let (negative, digits) = match parse_sign(trimmed_bytes) {
Some(result) => result,
None => return Ok(None),
};
let mut result: T = T::zero();
let radix = T::from(10_u8);
let stop_value = min_value / radix;
let mut iter = digits.iter();
// Parse integer portion until '.' or end
for &ch in iter.by_ref() {
if ch == b'.' {
break;
}
if !ch.is_ascii_digit() {
return Ok(None);
}
if result < stop_value {
return Ok(None);
}
let v = result * radix;
let digit: T = T::from(ch - b'0');
match v.checked_sub(&digit) {
Some(x) if x <= T::zero() => result = x,
_ => return Ok(None),
}
}
// Validate decimal portion (digits only, values ignored)
for &ch in iter {
if !ch.is_ascii_digit() {
return Ok(None);
}
}
Ok(finalize_int_result(result, negative))
}
fn do_parse_string_to_int_ansi<T: Integer + CheckedSub + CheckedNeg + From<u8> + Copy>(
str: &str,
type_name: &str,
min_value: T,
) -> SparkResult<Option<T>> {
let error = || Err(invalid_value(str, "STRING", type_name));
let trimmed_bytes = trim_all_bytes(str.as_bytes());
let (negative, digits) = match parse_sign(trimmed_bytes) {
Some(result) => result,
None => return error(),
};
let mut result: T = T::zero();
let radix = T::from(10_u8);
let stop_value = min_value / radix;
for &ch in digits {
if ch == b'.' || !ch.is_ascii_digit() {
return error();
}
if result < stop_value {
return error();
}
let v = result * radix;
let digit: T = T::from(ch - b'0');
match v.checked_sub(&digit) {
Some(x) if x <= T::zero() => result = x,
_ => return error(),
}
}
finalize_int_result(result, negative)
.map(Some)
.ok_or_else(|| invalid_value(str, "STRING", type_name))
}
fn do_parse_string_to_int_try<T: Integer + CheckedSub + CheckedNeg + From<u8> + Copy>(
str: &str,
min_value: T,
) -> SparkResult<Option<T>> {
let trimmed_bytes = trim_all_bytes(str.as_bytes());
let (negative, digits) = match parse_sign(trimmed_bytes) {
Some(result) => result,
None => return Ok(None),
};
let mut result: T = T::zero();
let radix = T::from(10_u8);
let stop_value = min_value / radix;
for &ch in digits {
if ch == b'.' || !ch.is_ascii_digit() {
return Ok(None);
}
if result < stop_value {
return Ok(None);
}
let v = result * radix;
let digit: T = T::from(ch - b'0');
match v.checked_sub(&digit) {
Some(x) if x <= T::zero() => result = x,
_ => return Ok(None),
}
}
Ok(finalize_int_result(result, negative))
}
fn parse_string_to_i8_legacy(str: &str) -> SparkResult<Option<i8>> {
match do_parse_string_to_int_legacy::<i32>(str, i32::MIN)? {
Some(v) if v >= i8::MIN as i32 && v <= i8::MAX as i32 => Ok(Some(v as i8)),
_ => Ok(None),
}
}
fn parse_string_to_i8_ansi(str: &str) -> SparkResult<Option<i8>> {
match do_parse_string_to_int_ansi::<i32>(str, "TINYINT", i32::MIN)? {
Some(v) if v >= i8::MIN as i32 && v <= i8::MAX as i32 => Ok(Some(v as i8)),
_ => Err(invalid_value(str, "STRING", "TINYINT")),
}
}
fn parse_string_to_i8_try(str: &str) -> SparkResult<Option<i8>> {
match do_parse_string_to_int_try::<i32>(str, i32::MIN)? {
Some(v) if v >= i8::MIN as i32 && v <= i8::MAX as i32 => Ok(Some(v as i8)),
_ => Ok(None),
}
}
fn parse_string_to_i16_legacy(str: &str) -> SparkResult<Option<i16>> {
match do_parse_string_to_int_legacy::<i32>(str, i32::MIN)? {
Some(v) if v >= i16::MIN as i32 && v <= i16::MAX as i32 => Ok(Some(v as i16)),
_ => Ok(None),
}
}
fn parse_string_to_i16_ansi(str: &str) -> SparkResult<Option<i16>> {
match do_parse_string_to_int_ansi::<i32>(str, "SMALLINT", i32::MIN)? {
Some(v) if v >= i16::MIN as i32 && v <= i16::MAX as i32 => Ok(Some(v as i16)),
_ => Err(invalid_value(str, "STRING", "SMALLINT")),
}
}
fn parse_string_to_i16_try(str: &str) -> SparkResult<Option<i16>> {
match do_parse_string_to_int_try::<i32>(str, i32::MIN)? {
Some(v) if v >= i16::MIN as i32 && v <= i16::MAX as i32 => Ok(Some(v as i16)),
_ => Ok(None),
}
}
/// Parses sign and returns (is_negative, remaining_bytes after sign)
/// Returns None if invalid (empty input, or just "+" or "-")
fn parse_sign(bytes: &[u8]) -> Option<(bool, &[u8])> {
let (&first, rest) = bytes.split_first()?;
match first {
b'-' if !rest.is_empty() => Some((true, rest)),
b'+' if !rest.is_empty() => Some((false, rest)),
_ => Some((false, bytes)),
}
}
#[inline]
pub fn invalid_value(value: &str, from_type: &str, to_type: &str) -> SparkError {
SparkError::CastInvalidValue {
value: value.to_string(),
from_type: from_type.to_string(),
to_type: to_type.to_string(),
}
}
fn parse_to_timestamp_info(
value: &str,
timestamp_type: &str,
) -> SparkResult<Option<TimeStampInfo>> {
let (sign, date_part) = if let Some(stripped) = value.strip_prefix('-') {
(-1i32, stripped)
} else {
(1i32, value)
};
let mut parts = date_part.split(['T', ' ', '-', ':', '.']);
let year = sign
* parts
.next()
.unwrap_or("")
.parse::<i32>()
.unwrap_or_default();
// Guard against years that cannot produce a valid i64 microsecond timestamp.
// The Long.MaxValue/MinValue boundaries correspond to years 294247 / -290308.
// We allow a slightly wider range and let parse_timestamp_to_micros perform the
// exact overflow check via checked arithmetic.
if !(-290309..=294248).contains(&year) {
return Ok(None);
}
let month = parts.next().map_or(1, |m| m.parse::<u32>().unwrap_or(1));
let day = parts.next().map_or(1, |d| d.parse::<u32>().unwrap_or(1));
let hour = parts.next().map_or(0, |h| h.parse::<u32>().unwrap_or(0));
let minute = parts.next().map_or(0, |m| m.parse::<u32>().unwrap_or(0));
let second = parts.next().map_or(0, |s| s.parse::<u32>().unwrap_or(0));
let microsecond = if let Some(ms) = parts.next() {
let Some(ms) = ms.get(..ms.len().min(6)) else {
return Ok(None);
};
let n = ms.len();
ms.parse::<u32>().unwrap_or(0) * 10u32.pow((6 - n) as u32)
} else {
0
};
let mut timestamp_info = TimeStampInfo::default();
let timestamp_info = match timestamp_type {
"year" => timestamp_info.with_year(year),
"month" => timestamp_info.with_year(year).with_month(month),
"day" => timestamp_info
.with_year(year)
.with_month(month)
.with_day(day),
"hour" => timestamp_info
.with_year(year)
.with_month(month)
.with_day(day)
.with_hour(hour),
"minute" => timestamp_info
.with_year(year)
.with_month(month)
.with_day(day)
.with_hour(hour)
.with_minute(minute),
"second" => timestamp_info
.with_year(year)
.with_month(month)
.with_day(day)
.with_hour(hour)
.with_minute(minute)
.with_second(second),
"microsecond" => timestamp_info
.with_year(year)
.with_month(month)
.with_day(day)
.with_hour(hour)
.with_minute(minute)
.with_second(second)
.with_microsecond(microsecond),
_ => {
return Err(SparkError::InvalidInputInCastToDatetime {
value: value.to_string(),
from_type: "STRING".to_string(),
to_type: "TIMESTAMP".to_string(),
})
}
};
Ok(Some(timestamp_info.to_owned()))
}
fn get_timestamp_values<T: TimeZone>(
value: &str,
timestamp_type: &str,
tz: &T,
) -> SparkResult<Option<i64>> {
match parse_to_timestamp_info(value, timestamp_type)? {
Some(info) => parse_timestamp_to_micros(&info, tz),
None => Ok(None),
}
}
/// Howard Hinnant's algorithm: proleptic Gregorian days since 1970-01-01 for any i64 year.
/// Works correctly for positive and negative years via Euclidean floor division.
/// Spark uses Java's equivalent [LocalDate.toEpochDay](https://github.com/openjdk/jdk/blob/cddee6d6eb3e048635c380a32bd2f6ebfd2c18b5/src/java.base/share/classes/java/time/LocalDate.java#L1954)
fn days_from_civil(y: i64, m: i64, d: i64) -> i64 {
let (y, m) = if m <= 2 { (y - 1, m + 9) } else { (y, m - 3) };
let era = if y >= 0 { y / 400 } else { (y - 399) / 400 };
let yoe = y - era * 400; // year of era [0, 399]
let doy = (153 * m + 2) / 5 + d - 1; // day of year [0, 365]
let doe = yoe * 365 + yoe / 4 - yoe / 100 + doy; // day of era [0, 146096]
era * 146097 + doe - 719468
}
fn is_leap_year(year: i64) -> bool {
year % 4 == 0 && (year % 100 != 0 || year % 400 == 0)
}
/// Days since 1970-01-01 for a proleptic Gregorian year/month/day, or `None` when the
/// combination is not a real calendar date. Unlike `NaiveDate::from_ymd_opt`, this accepts
/// any year that fits in `i64`.
pub(crate) fn ymd_to_epoch_day(year: i64, month: i64, day: i64) -> Option<i64> {
const DAYS_IN_MONTH: [i64; 12] = [31, 28, 31, 30, 31, 30, 31, 31, 30, 31, 30, 31];
let mut max_day = *DAYS_IN_MONTH.get(usize::try_from(month.checked_sub(1)?).ok()?)?;
if month == 2 && is_leap_year(year) {
max_day = 29;
}
if day < 1 || day > max_day {
return None;
}
Some(days_from_civil(year, month, day))
}
fn parse_timestamp_to_micros<T: TimeZone>(
timestamp_info: &TimeStampInfo,
tz: &T,
) -> SparkResult<Option<i64>> {
// Build NaiveDateTime explicitly so we can pattern-match LocalResult variants and
// handle the DST spring-forward gap case.
let naive_date_opt = NaiveDate::from_ymd_opt(
timestamp_info.year,
timestamp_info.month,
timestamp_info.day,
);
// NaiveTime is used for the common path; also validates hour/min/sec.
let naive_time = match NaiveTime::from_hms_opt(
timestamp_info.hour,
timestamp_info.minute,
timestamp_info.second,
) {
Some(t) => t,
None => return Ok(None), // invalid time components
};
if let Some(naive_date) = naive_date_opt {
let local_naive = naive_date.and_time(naive_time);
// Resolve local datetime to UTC, handling DST transitions.
// We compute base_micros with second precision (local_naive has no sub-second component),
// then add microseconds at the end to avoid calling with_nanosecond(), which internally
// calls from_local_datetime().single() and returns None for ambiguous (fall-back) times.
let base_micros: Option<i64> = match tz.from_local_datetime(&local_naive) {
// Unambiguous local time.
LocalResult::Single(dt) => Some(dt.timestamp_micros()),
// DST fall-back overlap: Spark picks the earlier UTC instant (pre-transition offset).
LocalResult::Ambiguous(earlier, _) => Some(earlier.timestamp_micros()),
// DST spring-forward gap: the local time does not exist.
// Java's ZonedDateTime.of() advances by the gap length, which is equivalent to
// utc = local_naive − pre_gap_offset
LocalResult::None => {
let probe = local_naive - chrono::Duration::hours(3);
let pre_offset = match tz.from_local_datetime(&probe) {
LocalResult::Single(dt) => dt.offset().fix(),
LocalResult::Ambiguous(dt, _) => dt.offset().fix(),
LocalResult::None => return Ok(None),
};
let offset_secs = pre_offset.local_minus_utc() as i64;
let utc_naive = local_naive - chrono::Duration::seconds(offset_secs);
Some(utc_naive.and_utc().timestamp_micros())
}
};
Ok(base_micros.map(|m| m + timestamp_info.microsecond as i64))
} else {
// NaiveDate::from_ymd_opt returned None. This means either:
// (a) invalid calendar date (Feb 29 on non-leap year, month 13, etc.)
// (b) year outside chrono's representable range (> 262143 or < -262144)
//
// For case (b) we fall back to Howard Hinnant's direct arithmetic, which works
// for any year that fits in i64. This covers the Long.MaxValue / Long.MinValue
// boundary timestamps (year 294247 / -290308).
let year = timestamp_info.year as i64;
if (-262144..=262143).contains(&year) {
// Year is in chrono's range but date was rejected -> truly invalid date.
return Ok(None);
}
// Validate month and day manually for extreme years.
let days =
match ymd_to_epoch_day(year, timestamp_info.month as i64, timestamp_info.day as i64) {
Some(days) => days,
None => return Ok(None),
};
// Compute the timezone offset using epoch as a surrogate probe point.
// Extreme-year timestamps are only valid with a UTC-like fixed offset (any DST
// zone would overflow). Using epoch gives us the standard offset.
let epoch_probe = NaiveDate::from_ymd_opt(1970, 1, 1)
.unwrap()
.and_hms_opt(0, 0, 0)
.unwrap();
let tz_offset_secs: i64 = match tz.from_local_datetime(&epoch_probe) {
LocalResult::Single(dt) => dt.offset().fix().local_minus_utc() as i64,
LocalResult::Ambiguous(dt, _) => dt.offset().fix().local_minus_utc() as i64,
LocalResult::None => 0,
};
// Compute seconds since epoch via direct calendar arithmetic.
// Use i128 for the intermediate multiply-by-1_000_000 step: the seconds value can be
// just outside the i64 range while the final microseconds result is still within range
// (e.g., Long.MinValue boundary: seconds = -9_223_372_036_855, result = i64::MIN).
let time_secs = timestamp_info.hour as i64 * 3600
+ timestamp_info.minute as i64 * 60
+ timestamp_info.second as i64;
let total_secs = days
.checked_mul(86400)
.and_then(|s| s.checked_add(time_secs))
.and_then(|s| s.checked_sub(tz_offset_secs));
let utc_micros = total_secs.and_then(|s| {
let micros128 = s as i128 * 1_000_000 + timestamp_info.microsecond as i128;
i64::try_from(micros128).ok()
});
Ok(utc_micros)
}
}
fn local_datetime_to_micros(timestamp_info: &TimeStampInfo) -> SparkResult<Option<i64>> {
let year = timestamp_info.year as i64;
let days = match ymd_to_epoch_day(year, timestamp_info.month as i64, timestamp_info.day as i64)
{
Some(days) => days,
None => return Ok(None),
};
if timestamp_info.hour >= 24 || timestamp_info.minute >= 60 || timestamp_info.second >= 60 {
return Ok(None);
}
let time_secs = timestamp_info.hour as i64 * 3600
+ timestamp_info.minute as i64 * 60
+ timestamp_info.second as i64;
let total_secs = days
.checked_mul(86400)
.and_then(|s| s.checked_add(time_secs));
let micros = total_secs.and_then(|s| {
let micros128 = s as i128 * 1_000_000 + timestamp_info.microsecond as i128;
i64::try_from(micros128).ok()
});
Ok(micros)
}
fn parse_str_to_year_timestamp<T: TimeZone>(value: &str, tz: &T) -> SparkResult<Option<i64>> {
get_timestamp_values(value, "year", tz)
}
fn parse_str_to_month_timestamp<T: TimeZone>(value: &str, tz: &T) -> SparkResult<Option<i64>> {
get_timestamp_values(value, "month", tz)
}
fn parse_str_to_day_timestamp<T: TimeZone>(value: &str, tz: &T) -> SparkResult<Option<i64>> {
get_timestamp_values(value, "day", tz)
}
fn parse_str_to_hour_timestamp<T: TimeZone>(value: &str, tz: &T) -> SparkResult<Option<i64>> {
get_timestamp_values(value, "hour", tz)
}
fn parse_str_to_minute_timestamp<T: TimeZone>(value: &str, tz: &T) -> SparkResult<Option<i64>> {
get_timestamp_values(value, "minute", tz)
}
fn parse_str_to_second_timestamp<T: TimeZone>(value: &str, tz: &T) -> SparkResult<Option<i64>> {
get_timestamp_values(value, "second", tz)
}
fn parse_str_to_microsecond_timestamp<T: TimeZone>(
value: &str,
tz: &T,
) -> SparkResult<Option<i64>> {
get_timestamp_values(value, "microsecond", tz)
}
fn timestamp_parser<T: TimeZone>(
value: &str,
eval_mode: EvalMode,
tz: &T,
is_spark4_plus: bool,
) -> SparkResult<Option<i64>> {
let trimmed = value.trim();
if trimmed.is_empty() {
return Ok(None);
}
// Spark 4.0+ rejects leading whitespace for ALL T-prefixed time-only strings
// (T<h>, T<h>:<m>, T<h>:<m>:<s>, T<h>:<m>:<s>.<f>), but accepts trailing whitespace.
// Spark 3.x trims all whitespace first, so leading whitespace is accepted there.
// Check the prefix, not the base patterns: a zone suffix can hide a time-only match.
if is_spark4_plus && value.len() > value.trim_start().len() && trimmed.starts_with('T') {
return if eval_mode == EvalMode::Ansi {
Err(SparkError::InvalidInputInCastToDatetime {
value: value.to_string(),
from_type: "STRING".to_string(),
to_type: "TIMESTAMP".to_string(),
})
} else {
Ok(None)
};
}
let value = trimmed;
// Spark accepts a leading '+' year sign on full date-time strings (e.g. "+2020-01-01T12:34:56")
// but rejects it on time-only strings (e.g. "+12:12:12" -> null).
// Detect: '+' followed by at least one digit and then a '-' separator -> year prefix -> strip '+'.
// Anything else starting with '+' (time-only, bare number, etc.) -> null.
let value = if let Some(rest) = value.strip_prefix('+') {
let first_non_digit = rest.find(|c: char| !c.is_ascii_digit());
match first_non_digit {
Some(i) if i >= 1 && rest.as_bytes()[i] == b'-' => rest,
_ => return Ok(None),
}
} else {
value
};
// Only attempt offset-suffix extraction when the value does not already match a
// base pattern. This prevents the '-' in plain date strings like "2015-03-18"
// from being misidentified as a negative-offset sign.
let has_direct_match = RE_YEAR.is_match(value)
|| RE_MONTH.is_match(value)
|| RE_DAY.is_match(value)
|| RE_HOUR.is_match(value)
|| RE_MINUTE.is_match(value)
|| RE_SECOND.is_match(value)
|| RE_MICROSECOND.is_match(value)
|| RE_TIME_ONLY_H.is_match(value)
|| RE_TIME_ONLY_HM.is_match(value)
|| RE_TIME_ONLY_HMS.is_match(value)
|| RE_TIME_ONLY_HMSU.is_match(value)
|| RE_BARE_HM.is_match(value)
|| RE_BARE_HMS.is_match(value)
|| RE_BARE_HMSU.is_match(value);
if !has_direct_match {
if let Some((stripped, suffix_tz)) = extract_offset_suffix(value) {
// Spark applies Java String.trim to the zone, not Unicode whitespace trimming.
let stripped = stripped.trim_end_matches(|c: char| c <= '\u{20}');
// A zone suffix is only meaningful after the seconds segment. Otherwise fall
// through with the unstripped value, which no base pattern matches, so it is
// reported as malformed (null, or CAST_INVALID_INPUT under ANSI) like Spark does.
if ends_with_seconds_segment(stripped) {
return timestamp_parser_with_tz(stripped, eval_mode, &suffix_tz);
}
}
}
timestamp_parser_with_tz(value, eval_mode, tz)
}
/// Parses the portion of an offset string AFTER any "UTC"/"GMT"/"UT" prefix (or the
/// full bare +/- offset including its sign character). Returns the offset in whole seconds,
/// or `None` for any malformed, out-of-range, or trailing-garbage input.
///
/// Accepted formats (H = 1–2 digit hour, M = 1–2 digit minute):
/// "" -> 0 (bare "UTC" / "GMT" / "UT")
/// "+H" -> +H*3600 (hour-only, e.g. "+0" from "UTC+0")
/// "+HH" -> same
/// "+HHMM" -> +H*3600+M*60 (4 digits, no colon)
/// "+H:M" -> same (with colon, any digit count 1-2 each)
/// "+HH:MM" -> same
/// (negative with '-' analogously)
///
/// Hours must be 0–18 and minutes 0–59, with a maximum absolute offset of 18:00.
/// A trailing colon ("+8:") is rejected.
fn parse_sign_offset(s: &str) -> Option<i32> {
if s.is_empty() {
return Some(0);
}
let (sign, rest) = match s.as_bytes().first() {
Some(&b'+') => (1i32, &s[1..]),
Some(&b'-') => (-1i32, &s[1..]),
_ => return None,
};
// Validate before slicing: malformed date segments can reach this helper, and a
// byte range such as rest[..2] must not split a non-ASCII digit's UTF-8 encoding.
if rest.is_empty() || !rest.bytes().all(|b| b.is_ascii_digit() || b == b':') {
return None;
}
let (h, m) = if let Some(colon_pos) = rest.find(':') {
let h_str = &rest[..colon_pos];
let m_str = &rest[colon_pos + 1..];
if !(1..=2).contains(&h_str.len()) || !(1..=2).contains(&m_str.len()) {
return None;
}
let h: i32 = h_str.parse().ok()?;
// Note: "+HH:MM:SS" (with seconds) is not handled; Spark accepts it but it is rare.
let m: i32 = m_str.parse().ok()?;
(h, m)
} else {
match rest.len() {
1 | 2 => (rest.parse::<i32>().ok()?, 0),
4 => (
rest[..2].parse::<i32>().ok()?,
rest[2..].parse::<i32>().ok()?,
),
_ => return None,
}
};
if !(0..=18).contains(&h) || !(0..=59).contains(&m) || (h == 18 && m != 0) {
return None;
}
Some(sign * (h * 3600 + m * 60))
}
/// Constructs a [`Tz`] from an offset measured in seconds.
/// E.g. `+7*3600 + 30*60` -> `"+07:30"`.
fn tz_from_offset_secs(secs: i32) -> Option<Tz> {
let abs = secs.abs();
let h = abs / 3600;
let m = (abs % 3600) / 60;
let sign = if secs >= 0 { '+' } else { '-' };
Tz::from_str(&format!("{}{:02}:{:02}", sign, h, m)).ok()
}
/// Returns the last (rightmost) byte position where `needle` starts inside `haystack`.
fn rfind_str(haystack: &str, needle: &str) -> Option<usize> {
let hb = haystack.as_bytes();
let nb = needle.as_bytes();
if nb.len() > hb.len() {
return None;
}
(0..=(hb.len() - nb.len()))
.rev()
.find(|&i| hb[i..].starts_with(nb))
}
/// If `value` ends with a recognised timezone suffix, returns `(datetime_prefix, Tz)`.
/// Returns `None` when no suffix is found.
///
/// Recognised forms (in matching priority order):
/// Z -> UTC+0
/// UTC / " UTC" -> UTC+0 (or UTC +/- offset, e.g. "UTC+0", " UTC+07:30")
/// GMT / " GMT" -> UTC+0 (or GMT +/- offset)
/// UT / " UT" -> UTC+0 (or UT +/- offset)
/// Named IANA zone -> e.g. " Europe/Moscow"
/// Bare +/-offset -> e.g. "+07:30", "-1:0", "+0730"
///
/// **The caller must ensure the value does not already match a base timestamp pattern.**
/// Without that guard a bare '-' in "2015-03-18" would be misread as a -18:00 offset.
fn extract_offset_suffix(value: &str) -> Option<(&str, Tz)> {
// 1. Z suffix
if let Some(stripped) = value.strip_suffix('Z') {
return Some((stripped, tz_from_offset_secs(0)?));
}
// 2. Named text-prefix forms: "UTC", "GMT", "UT" (optionally space-prefixed),
// each optionally followed by a bare +/-offset.
// Longest first so " UTC" is tried before " UT", etc.
for prefix in &[" UTC", "UTC", " GMT", "GMT", " UT", "UT"] {
if let Some(pos) = rfind_str(value, prefix) {
let offset_str = &value[pos + prefix.len()..];
if let Some(secs) = parse_sign_offset(offset_str) {
return Some((&value[..pos], tz_from_offset_secs(secs)?));
}
}
}
// 3. Java SHORT_IDS fixed-offset abbreviations recognised by ZoneId.of() via SHORT_IDS map.
// Only three have purely fixed offsets (no '/'):
// EST -> -05:00 (-18 000 s)
// MST -> -07:00 (-25 200 s)
// HST -> -10:00 (-36 000 s)
// These may appear with or without a leading space; no sub-offset is allowed after them.
for (abbr, offset_secs) in &[
(" EST", -18_000i32),
("EST", -18_000),
(" MST", -25_200),
("MST", -25_200),
(" HST", -36_000),
("HST", -36_000),
] {
if let Some(pos) = rfind_str(value, abbr) {
if pos + abbr.len() == value.len() {
return Some((&value[..pos], tz_from_offset_secs(*offset_secs)?));
}
}
}
// 4. Named IANA timezone: a space followed by a slash-containing word at the end.
// e.g. "2015-03-18T12:03:17.123456 Europe/Moscow"
if let Some(space_pos) = value.rfind(' ') {
let tz_name = &value[space_pos + 1..];
if tz_name.contains('/') {
if let Ok(tz) = Tz::from_str(tz_name) {
return Some((&value[..space_pos], tz));
}
}
}
// 5. Bare +/-offset: find the rightmost '+' or '-' and try to parse everything
// from that position to the end as a complete valid offset.
let last_sign = {
let p = value.rfind('+');
let m = value.rfind('-');
match (p, m) {
(Some(p), Some(m)) => Some(p.max(m)),
(a, b) => a.or(b),
}
};
if let Some(pos) = last_sign {
let offset_str = &value[pos..];
if let Some(secs) = parse_sign_offset(offset_str) {
return Some((&value[..pos], tz_from_offset_secs(secs)?));
}
}
None
}
type TimestampParsePattern<T> = (&'static Regex, fn(&str, &T) -> SparkResult<Option<i64>>);
// These shapes transcribe the per-segment digit rules of Spark's
// `SparkDateTimeUtils.parseTimestampString` (`isValidDigits`): the year takes 4-6 digits
// (`maxDigitsYear = 6`, so "0002020-01-01" is malformed for a timestamp even though
// `stringToDate`, ported by `date_parser`, allows 7), month/day/hour/minute/second take 1-2
// digits each, and the fraction takes any number of digits including none ("12:34:56." is
// valid), of which only the first six are kept. All digits must be ASCII, matching Spark's
// byte scanner and the numeric parsers used after shape recognition.
// Keep the ASCII ranges: Unicode `\d` also costs substantially more to match on valid input.
static RE_YEAR: LazyLock<Regex> = LazyLock::new(|| Regex::new(r"^-?[0-9]{4,6}$").unwrap());
static RE_MONTH: LazyLock<Regex> =
LazyLock::new(|| Regex::new(r"^-?[0-9]{4,6}-[0-9]{1,2}$").unwrap());
static RE_DAY: LazyLock<Regex> =
LazyLock::new(|| Regex::new(r"^-?[0-9]{4,6}-[0-9]{1,2}-[0-9]{1,2}$").unwrap());
static RE_HOUR: LazyLock<Regex> =
LazyLock::new(|| Regex::new(r"^-?[0-9]{4,6}-[0-9]{1,2}-[0-9]{1,2}[T ][0-9]{1,2}$").unwrap());
static RE_MINUTE: LazyLock<Regex> = LazyLock::new(|| {
Regex::new(r"^-?[0-9]{4,6}-[0-9]{1,2}-[0-9]{1,2}[T ][0-9]{1,2}:[0-9]{1,2}$").unwrap()
});
static RE_SECOND: LazyLock<Regex> = LazyLock::new(|| {
Regex::new(r"^-?[0-9]{4,6}-[0-9]{1,2}-[0-9]{1,2}[T ][0-9]{1,2}:[0-9]{1,2}:[0-9]{1,2}$").unwrap()
});
static RE_MICROSECOND: LazyLock<Regex> = LazyLock::new(|| {
Regex::new(r"^-?[0-9]{4,6}-[0-9]{1,2}-[0-9]{1,2}[T ][0-9]{1,2}:[0-9]{1,2}:[0-9]{1,2}\.[0-9]*$")
.unwrap()
});
static RE_TIME_ONLY_H: LazyLock<Regex> = LazyLock::new(|| Regex::new(r"^T[0-9]{1,2}$").unwrap());
static RE_TIME_ONLY_HM: LazyLock<Regex> =
LazyLock::new(|| Regex::new(r"^T[0-9]{1,2}:[0-9]{1,2}$").unwrap());
static RE_TIME_ONLY_HMS: LazyLock<Regex> =
LazyLock::new(|| Regex::new(r"^T[0-9]{1,2}:[0-9]{1,2}:[0-9]{1,2}$").unwrap());
static RE_TIME_ONLY_HMSU: LazyLock<Regex> =
LazyLock::new(|| Regex::new(r"^T[0-9]{1,2}:[0-9]{1,2}:[0-9]{1,2}\.[0-9]*$").unwrap());
static RE_BARE_HM: LazyLock<Regex> =
LazyLock::new(|| Regex::new(r"^[0-9]{1,2}:[0-9]{1,2}$").unwrap());
static RE_BARE_HMS: LazyLock<Regex> =
LazyLock::new(|| Regex::new(r"^[0-9]{1,2}:[0-9]{1,2}:[0-9]{1,2}$").unwrap());
static RE_BARE_HMSU: LazyLock<Regex> =
LazyLock::new(|| Regex::new(r"^[0-9]{1,2}:[0-9]{1,2}:[0-9]{1,2}\.[0-9]*$").unwrap());
/// Whether `value` (a datetime with any zone suffix already stripped) ends in a seconds or
/// fraction segment. Spark's `parseTimestampString` only captures a zone id when its byte
/// scanner hits a non-digit while inside those two segments, so a suffix such as `Z`, `+05:30`
/// or ` UTC` is legal after `hh:mm:ss` or `hh:mm:ss.f*` but makes a date-only, hour-only or
/// hour:minute value malformed ("2020-10-01Z" and "2020-01-01T12:34Z" are both null).
fn ends_with_seconds_segment(value: &str) -> bool {
RE_SECOND.is_match(value)
|| RE_MICROSECOND.is_match(value)
|| RE_TIME_ONLY_HMS.is_match(value)
|| RE_TIME_ONLY_HMSU.is_match(value)
|| RE_BARE_HMS.is_match(value)
|| RE_BARE_HMSU.is_match(value)
}
fn timestamp_parser_with_tz<T: TimeZone>(
value: &str,
eval_mode: EvalMode,
tz: &T,
) -> SparkResult<Option<i64>> {
// Both T-separator and space-separator date-time forms are supported.
// Negative years are handled by get_timestamp_values detecting a leading '-'.
let patterns: &[TimestampParsePattern<T>] = &[
// Year only: 4-6 digits, optionally negative
(
&RE_YEAR,
parse_str_to_year_timestamp as fn(&str, &T) -> SparkResult<Option<i64>>,
),
// Year-month
(&RE_MONTH, parse_str_to_month_timestamp),
// Year-month-day
(&RE_DAY, parse_str_to_day_timestamp),
// Date T-or-space hour (1 or 2 digits)
(&RE_HOUR, parse_str_to_hour_timestamp),
// Date T-or-space hour:minute
(&RE_MINUTE, parse_str_to_minute_timestamp),
// Date T-or-space hour:minute:second
(&RE_SECOND, parse_str_to_second_timestamp),
// Date T-or-space hour:minute:second.fraction
(&RE_MICROSECOND, parse_str_to_microsecond_timestamp),
// Time-only: T hour (1 or 2 digits, no colon)
(&RE_TIME_ONLY_H, parse_str_to_time_only_timestamp),
// Time-only: T hour:minute
(&RE_TIME_ONLY_HM, parse_str_to_time_only_timestamp),
// Time-only: T hour:minute:second
(&RE_TIME_ONLY_HMS, parse_str_to_time_only_timestamp),
// Time-only: T hour:minute:second.fraction
(&RE_TIME_ONLY_HMSU, parse_str_to_time_only_timestamp),
// Bare time-only: hour:minute (without T prefix)
(&RE_BARE_HM, parse_str_to_time_only_timestamp),
// Bare time-only: hour:minute:second
(&RE_BARE_HMS, parse_str_to_time_only_timestamp),
// Bare time-only: hour:minute:second.fraction
(&RE_BARE_HMSU, parse_str_to_time_only_timestamp),
];
let mut timestamp = None;
// Iterate through patterns and try matching
for (pattern, parse_func) in patterns {
if pattern.is_match(value) {
timestamp = parse_func(value, tz)?;
break;
}
}
if timestamp.is_none() {
return if eval_mode == EvalMode::Ansi {
Err(SparkError::InvalidInputInCastToDatetime {
value: value.to_string(),
from_type: "STRING".to_string(),
to_type: "TIMESTAMP".to_string(),
})
} else {
Ok(None)
};
}
Ok(timestamp)
}
fn timestamp_ntz_parser(
value: &str,
eval_mode: EvalMode,
allow_time_zone: bool,
_is_spark4_plus: bool,
) -> SparkResult<Option<i64>> {
let trimmed = value.trim();
if trimmed.is_empty() {
return Ok(None);
}
// NTZ rejects leading whitespace for T-prefixed time-only strings on Spark 4+
// (same logic as timestamp_parser), but time-only is rejected entirely for NTZ anyway.
let value = trimmed;
// Handle leading '+' the same way as timestamp_parser
let value = if let Some(rest) = value.strip_prefix('+') {
let first_non_digit = rest.find(|c: char| !c.is_ascii_digit());
match first_non_digit {
Some(i) if i >= 1 && rest.as_bytes()[i] == b'-' => rest,
_ => return Ok(None),
}
} else {
value
};
// Reject time-only patterns: NTZ requires a date component
if RE_TIME_ONLY_H.is_match(value)
|| RE_TIME_ONLY_HM.is_match(value)
|| RE_TIME_ONLY_HMS.is_match(value)
|| RE_TIME_ONLY_HMSU.is_match(value)
|| RE_BARE_HM.is_match(value)
|| RE_BARE_HMS.is_match(value)
|| RE_BARE_HMSU.is_match(value)
{
return if eval_mode == EvalMode::Ansi {
Err(SparkError::InvalidInputInCastToDatetime {
value: value.to_string(),
from_type: "STRING".to_string(),
to_type: "TIMESTAMP_NTZ".to_string(),
})
} else {
Ok(None)
};
}
// Check if value matches a date-based pattern directly
let has_direct_match = RE_YEAR.is_match(value)
|| RE_MONTH.is_match(value)
|| RE_DAY.is_match(value)
|| RE_HOUR.is_match(value)
|| RE_MINUTE.is_match(value)
|| RE_SECOND.is_match(value)
|| RE_MICROSECOND.is_match(value);
// If no direct match, try stripping a timezone suffix. Spark only recognises a zone after
// the seconds segment; a suffix anywhere else leaves the unstripped value, which no base
// pattern matches, so the inner parser reports it as malformed.
let value_to_parse = if !has_direct_match {
match extract_offset_suffix(value) {
Some((stripped, _tz))
if ends_with_seconds_segment(
stripped.trim_end_matches(|c: char| c <= '\u{20}'),
) =>
{
if !allow_time_zone {
return if eval_mode == EvalMode::Ansi {
Err(SparkError::InvalidInputInCastToDatetime {
value: value.to_string(),
from_type: "STRING".to_string(),
to_type: "TIMESTAMP_NTZ".to_string(),
})
} else {
Ok(None)
};
}
stripped.trim_end_matches(|c: char| c <= '\u{20}')
}
_ => value,
}
} else {
value
};
timestamp_ntz_parser_inner(value_to_parse, eval_mode)
}
fn timestamp_ntz_parser_inner(value: &str, eval_mode: EvalMode) -> SparkResult<Option<i64>> {
let patterns: &[(&Regex, &str)] = &[
(&RE_YEAR, "year"),
(&RE_MONTH, "month"),
(&RE_DAY, "day"),
(&RE_HOUR, "hour"),
(&RE_MINUTE, "minute"),
(&RE_SECOND, "second"),
(&RE_MICROSECOND, "microsecond"),
];
for (re, ts_type) in patterns {
if re.is_match(value) {
if let Some(info) = parse_to_timestamp_info(value, ts_type)? {
if let Some(timestamp) = local_datetime_to_micros(&info)? {
return Ok(Some(timestamp));
}
}
break;
}
}
if eval_mode == EvalMode::Ansi {
Err(SparkError::InvalidInputInCastToDatetime {
value: value.to_string(),
from_type: "STRING".to_string(),
to_type: "TIMESTAMP_NTZ".to_string(),
})
} else {
Ok(None)
}
}
fn parse_str_to_time_only_timestamp<T: TimeZone>(value: &str, tz: &T) -> SparkResult<Option<i64>> {
// The 'T' is optional in the time format; strip it if specified.
let time_part = value.strip_prefix('T').unwrap_or(value);
// Parse time components: hour[:minute[:second[.fraction]]]
// Use splitn(3) so "12:34:56.789" splits into ["12", "34", "56.789"].
let colon_parts: Vec<&str> = time_part.splitn(3, ':').collect();
let hour: u32 = colon_parts[0].parse().unwrap_or(0);
let minute: u32 = colon_parts.get(1).and_then(|s| s.parse().ok()).unwrap_or(0);
let (second, nanosecond) = if let Some(sec_frac) = colon_parts.get(2) {
let dot_idx = sec_frac.find('.');
let sec: u32 = sec_frac[..dot_idx.unwrap_or(sec_frac.len())]
.parse()
.unwrap_or(0);
let ns: u32 = if let Some(dot) = dot_idx {
let frac = &sec_frac[dot + 1..];
// Interpret up to 6 digits as microseconds, padding with trailing zeros.
let Some(trimmed) = frac.get(..frac.len().min(6)) else {
return Ok(None);
};
let padded = format!("{:0<6}", trimmed);
padded.parse::<u32>().unwrap_or(0) * 1000
} else {
0
};
(sec, ns)
} else {
(0, 0)
};
let datetime = tz.from_utc_datetime(&chrono::Utc::now().naive_utc());
let result = datetime
.with_timezone(tz)
.with_hour(hour)
.and_then(|dt| dt.with_minute(minute))
.and_then(|dt| dt.with_second(second))
.and_then(|dt| dt.with_nanosecond(nanosecond))
.map(|dt| dt.timestamp_micros());
Ok(result)
}
//a string to date parser - port of spark's SparkDateTimeUtils#stringToDate.
fn date_parser(date_str: &str, eval_mode: EvalMode) -> SparkResult<Option<i32>> {
// local functions
/// Decodes a run of ASCII digits, or `None` if any byte is not a digit.
fn decode_digits(bytes: &[u8]) -> Option<i64> {
bytes.iter().try_fold(0i64, |acc, b| {
b.is_ascii_digit().then(|| acc * 10 + (b - b'0') as i64)
})
}
fn is_valid_digits(segment: i32, digits: usize) -> bool {
// Years are bounded by `resolve_epoch_day` below. We allow up to 7 digits to support
// leading-zero year strings like "0002020" (= year 2020), matching Spark's isValidDigits.
let max_digits_year = 7;
// year (segment 0) can be between 4 to 7 digits,
// month and day (segment 1 and 2) can be between 1 to 2 digits
(segment == 0 && digits >= 4 && digits <= max_digits_year)
|| (segment != 0 && digits > 0 && digits <= 2)
}
fn return_result(date_str: &str, eval_mode: EvalMode) -> SparkResult<Option<i32>> {
if eval_mode == EvalMode::Ansi {
Err(SparkError::InvalidInputInCastToDatetime {
value: date_str.to_string(),
from_type: "STRING".to_string(),
to_type: "DATE".to_string(),
})
} else {
Ok(None)
}
}
/// Turns parsed year/month/day segments into an epoch day. Shared by both parsing paths so
/// that the decision of what a segment triple *means* lives in exactly one place.
fn resolve_epoch_day(
year: i64,
month: i64,
day: i64,
date_str: &str,
eval_mode: EvalMode,
) -> SparkResult<Option<i32>> {
// Spark builds a `LocalDate` and narrows its epoch day to an `Int`, so an invalid
// calendar date or an epoch day that overflows `i32` both yield `None` there, which
// `stringToDateAnsi` turns into CAST_INVALID_INPUT.
let Some(days) = ymd_to_epoch_day(year, month, day).and_then(|d| i32::try_from(d).ok())
else {
return return_result(date_str, eval_mode);
};
// Spark accepts years beyond what chrono can represent, and downstream date kernels
// cannot handle those values, so Comet keeps returning null for them in every eval mode
// rather than raising. This is a Comet limitation, not a malformed input.
//
// The bound is chrono's representable year range: `NaiveDate::MIN` is `-262143-01-01`
// and `NaiveDate::MAX` is `262142-12-31`
// (https://docs.rs/chrono/latest/chrono/naive/struct.NaiveDate.html#associatedconstant.MIN).
if !(-262143..=262142).contains(&year) {
return Ok(None);
}
Ok(Some(days))
}
// end local functions
if date_str.is_empty() {
return return_result(date_str, eval_mode);
}
let bytes = date_str.as_bytes();
// `SparkDateTimeUtils.getTrimmedStart`/`getTrimmedEnd` trim the same byte set as
// `UTF8String.trimAll`.
let (start, str_end_trimmed) = trim_all_range(bytes);
let mut j = start;
if j == str_end_trimmed {
return return_result(date_str, eval_mode);
}
// Fast path for the canonical `yyyy-mm-dd` form, which only skips the byte-scanning loop
// below. Any other shape (including a leading sign, which makes the first byte a
// non-digit) falls through to the general parser.
let trimmed = &bytes[j..str_end_trimmed];
if trimmed.len() == 10 && trimmed[4] == b'-' && trimmed[7] == b'-' {
if let (Some(year), Some(month), Some(day)) = (
decode_digits(&trimmed[..4]),
decode_digits(&trimmed[5..7]),
decode_digits(&trimmed[8..10]),
) {
return resolve_epoch_day(year, month, day, date_str, eval_mode);
}
}
//values of date segments year, month and day defaulting to 1
let mut date_segments = [1, 1, 1];
let mut sign = 1;
let mut current_segment = 0;
let mut current_segment_value = Wrapping(0);
let mut current_segment_digits = 0;
// assign a sign to the date; both '-' and '+' are accepted (Spark stringToDate line 357-360)
if bytes[j] == b'-' {
sign = -1;
j += 1;
} else if bytes[j] == b'+' {
// sign remains 1 (positive)
j += 1;
}
//loop to the end of string until we have processed 3 segments,
//exit loop on encountering any space ' ' or 'T' after the 3rd segment
while j < str_end_trimmed && (current_segment < 3 && !(bytes[j] == b' ' || bytes[j] == b'T')) {
let b = bytes[j];
if current_segment < 2 && b == b'-' {
//check for validity of year and month segments if current byte is separator
if !is_valid_digits(current_segment, current_segment_digits) {
return return_result(date_str, eval_mode);
}
//if valid update corresponding segment with the current segment value.
date_segments[current_segment as usize] = current_segment_value.0;
current_segment_value = Wrapping(0);
current_segment_digits = 0;
current_segment += 1;
} else if !b.is_ascii_digit() {
return return_result(date_str, eval_mode);
} else {
//increment value of current segment by the next digit
let parsed_value = Wrapping((b - b'0') as i32);
current_segment_value = current_segment_value * Wrapping(10) + parsed_value;
current_segment_digits += 1;
}
j += 1;
}
//check for validity of last segment
if !is_valid_digits(current_segment, current_segment_digits) {
return return_result(date_str, eval_mode);
}
if current_segment < 2 && j < str_end_trimmed {
// For the `yyyy` and `yyyy-[m]m` formats, entire input must be consumed.
return return_result(date_str, eval_mode);
}
date_segments[current_segment as usize] = current_segment_value.0;
resolve_epoch_day(
(sign * date_segments[0]) as i64,
date_segments[1] as i64,
date_segments[2] as i64,
date_str,
eval_mode,
)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::cast::cast_array;
use crate::SparkCastOptions;
use arrow::array::{DictionaryArray, Int32Array, StringArray};
use arrow::datatypes::TimeUnit;
use datafusion::common::Result as DataFusionResult;
/// Test helper that wraps the mode-specific parse functions
fn cast_string_to_i8(str: &str, eval_mode: EvalMode) -> SparkResult<Option<i8>> {
match eval_mode {
EvalMode::Legacy => parse_string_to_i8_legacy(str),
EvalMode::Ansi => parse_string_to_i8_ansi(str),
EvalMode::Try => parse_string_to_i8_try(str),
}
}
#[test]
fn test_digits_to_i128_boundary() {
// 38 nines is under i128::MAX (~1.7e38 vs 9.99e37); the head-only path parses it.
let d38 = "9".repeat(38);
assert_eq!(
digits_to_i128(d38.as_bytes()),
Some(99_999_999_999_999_999_999_999_999_999_999_999_999_i128)
);
// 39 nines overflows i128 in the tail-checked step.
let d39 = "9".repeat(39);
assert_eq!(digits_to_i128(d39.as_bytes()), None);
// 40 characters that reduce to a small value once the leading zero(s) are seen: the
// head-only accumulator must not lose bits on 38 zeros followed by two digits.
let z38_plus = format!("{}42", "0".repeat(38));
assert_eq!(digits_to_i128(z38_plus.as_bytes()), Some(42));
}
#[test]
fn test_parse_string_to_decimal_boundary() {
// 38-digit integral parses (fits i128).
let s38 = "9".repeat(38);
assert!(parse_string_to_decimal(&s38, 38, 0).unwrap().is_some());
// 39-digit integral overflows i128, so returns the invalid_decimal_cast error.
let s39 = "9".repeat(39);
assert!(parse_string_to_decimal(&s39, 38, 0).is_err());
// Very long fractional part now returns Err via the invalid_decimal_cast path instead
// of panicking on 10_i128.pow(fractional_scale).
let over_long = format!("0.{}", "0".repeat(40));
assert!(parse_string_to_decimal(&over_long, 38, 10).is_err());
}
#[test]
#[cfg_attr(miri, ignore)] // test takes too long with miri
fn test_cast_string_to_timestamp() {
let array: ArrayRef = Arc::new(StringArray::from(vec![
Some("2020-01-01T12:34:56.123456"),
Some("T2"),
Some("0100-01-01T12:34:56.123456"),
Some("10000-01-01T12:34:56.123456"),
// 7-digit year-only strings must return null (Spark returns null for these)
Some("0119704"),
Some("2024001"),
]));
let tz = &Tz::from_str("UTC").unwrap();
let string_array = array
.as_any()
.downcast_ref::<GenericStringArray<i32>>()
.expect("Expected a string array");
let eval_mode = EvalMode::Legacy;
let result = cast_utf8_to_timestamp!(
&string_array,
eval_mode,
PrimitiveArray::<TimestampMicrosecondType>::builder(string_array.len())
.with_timezone("UTC"),
timestamp_parser,
tz,
true
)
.unwrap();
assert_eq!(
result.data_type(),
&DataType::Timestamp(TimeUnit::Microsecond, Some("UTC".into()))
);
assert_eq!(result.len(), 6);
// 7-digit year-only strings must be null
assert!(result.is_null(4), "0119704 should be null");
assert!(result.is_null(5), "2024001 should be null");
}
#[test]
fn test_cast_string_to_timestamp_ansi_error() {
// In ANSI mode, an invalid timestamp string must produce an error rather than null.
let array: ArrayRef = Arc::new(StringArray::from(vec![
Some("2020-01-01T12:34:56.123456"),
Some("not_a_timestamp"),
]));
let tz = &Tz::from_str("UTC").unwrap();
let string_array = array
.as_any()
.downcast_ref::<GenericStringArray<i32>>()
.expect("Expected a string array");
let eval_mode = EvalMode::Ansi;
let result = cast_utf8_to_timestamp!(
&string_array,
eval_mode,
PrimitiveArray::<TimestampMicrosecondType>::builder(string_array.len())
.with_timezone("UTC"),
timestamp_parser,
tz,
true
);
assert!(
result.is_err(),
"ANSI mode should return Err for an invalid timestamp string"
);
}
#[test]
fn test_cast_string_to_timestamp_ansi_error_trimmed_value() {
// The error value in InvalidInputInCastToDatetime must match the raw input
// (including trailing whitespace) to match Spark's CAST_INVALID_INPUT behavior.
let array: ArrayRef = Arc::new(StringArray::from(vec![
Some("91\n3 "), // trailing spaces after a newline in the middle
]));
let tz = &Tz::from_str("UTC").unwrap();
let string_array = array
.as_any()
.downcast_ref::<GenericStringArray<i32>>()
.expect("Expected a string array");
let eval_mode = EvalMode::Ansi;
let result = cast_utf8_to_timestamp!(
&string_array,
eval_mode,
PrimitiveArray::<TimestampMicrosecondType>::builder(string_array.len())
.with_timezone("UTC"),
timestamp_parser,
tz,
true
);
match result {
Err(SparkError::InvalidInputInCastToDatetime { value, .. }) => {
assert_eq!(
value, "91\n3 ",
"ANSI error value should match the raw (untrimmed) input to match Spark behavior"
);
}
other => panic!("Expected InvalidInputInCastToDatetime error, got {other:?}"),
}
}
/// The codepoint matrix from
/// <https://github.com/apache/datafusion-comet/issues/5149>: the ASCII control bytes and
/// DELETE, plus the non-ASCII codepoints that are whitespace to Unicode but that no Spark
/// cast trims. `CometNativeCastSuite` runs the same matrix with Spark itself as the oracle.
fn trim_pads() -> Vec<String> {
let mut pads: Vec<String> = (0x00u8..=0x20).map(|b| String::from(b as char)).collect();
pads.push("\u{7f}".to_string());
pads.extend(
[
"\u{85}", "\u{a0}", "\u{1680}", "\u{2000}", "\u{2005}", "\u{200a}", "\u{2028}",
"\u{2029}", "\u{202f}", "\u{205f}", "\u{3000}",
]
.iter()
.map(|s| s.to_string()),
);
pads
}
/// Asserts that the cast to `to_type` trims each [`trim_pads`] entry exactly when `regime`
/// -- the trim helper that Spark's cast to `to_type` uses -- trims it: a value in every eval
/// mode when it is trimmed, NULL (or an ANSI error) when it is not. Interior padding, padding
/// on its own and the empty string must never parse.
fn assert_trim_parity(to_type: &DataType, valid: &str, regime: fn(&str) -> &str) {
let split = valid.char_indices().nth(1).map(|(i, _)| i).unwrap();
// The empty string reaches the same empty-slice branch that fully-trimmed padding does.
let mut cases = vec![("empty".to_string(), String::new(), false)];
for pad in trim_pads() {
// The regime trims this padding iff trimming the padding alone leaves nothing.
let trimmed = regime(&pad).is_empty();
cases.extend([
(format!("leading {pad:?}"), format!("{pad}{valid}"), trimmed),
(
format!("trailing {pad:?}"),
format!("{valid}{pad}"),
trimmed,
),
(
format!("both ends {pad:?}"),
format!("{pad}{valid}{pad}"),
trimmed,
),
(
format!("interior {pad:?}"),
format!("{}{pad}{}", &valid[..split], &valid[split..]),
false,
),
(format!("only {pad:?}"), pad.clone(), false),
]);
}
for (position, input, expect_value) in cases {
for eval_mode in [EvalMode::Legacy, EvalMode::Try, EvalMode::Ansi] {
let array: ArrayRef = Arc::new(StringArray::from(vec![Some(input.as_str())]));
let options = SparkCastOptions::new(eval_mode, "UTC", false);
let result = cast_array(array, to_type, &options);
let context = format!("cast {input:?} ({position}) to {to_type} in {eval_mode:?}");
if expect_value {
let array = result.unwrap_or_else(|e| panic!("{context}: {e}"));
assert!(!array.is_null(0), "{context}: expected a value, got NULL");
} else if eval_mode == EvalMode::Ansi {
assert!(result.is_err(), "{context}: expected an error");
} else {
let array = result.unwrap_or_else(|e| panic!("{context}: {e}"));
assert!(array.is_null(0), "{context}: expected NULL");
}
}
}
}
#[test]
fn test_cast_string_to_boolean_trim_parity() {
assert_trim_parity(&DataType::Boolean, "true", trim_all);
}
#[test]
fn test_cast_string_to_int_trim_parity() {
for to_type in [
DataType::Int8,
DataType::Int16,
DataType::Int32,
DataType::Int64,
] {
assert_trim_parity(&to_type, "12", trim_all);
}
}
#[test]
fn test_cast_string_to_date_trim_parity() {
assert_trim_parity(&DataType::Date32, "2020-01-01", trim_all);
}
/// Float, double and decimal all use the narrower `String.trim` set, which keeps `0x7F`.
#[test]
fn test_cast_string_to_float_and_decimal_trim_parity() {
for to_type in [
DataType::Float32,
DataType::Float64,
DataType::Decimal128(10, 2),
] {
assert_trim_parity(&to_type, "1.5", trim_java_string);
}
}
/// Pins the one trim divergence this PR leaves behind, so that resolving
/// <https://github.com/apache/datafusion-comet/issues/5149> has to update this test rather
/// than change behaviour silently. `timestamp_parser` and `timestamp_ntz_parser` still use
/// `str::trim`, so they accept the non-ASCII whitespace that Spark's
/// `SparkDateTimeUtils.getTrimmedStart` / `getTrimmedEnd` leave in place, where Spark returns
/// NULL. `CometNativeCastSuite` cannot cover this, because Spark is the oracle there and Comet does
/// not fall back -- it silently returns a value.
#[test]
fn test_cast_string_to_timestamp_unicode_whitespace_divergence() {
let to_types = [
DataType::Timestamp(TimeUnit::Microsecond, Some("UTC".into())),
DataType::Timestamp(TimeUnit::Microsecond, None),
];
for pad in ["\u{85}", "\u{a0}", "\u{2028}", "\u{3000}"] {
for to_type in &to_types {
let input = format!("{pad}2020-01-01 12:34:56{pad}");
let array: ArrayRef = Arc::new(StringArray::from(vec![Some(input.as_str())]));
let options = SparkCastOptions::new(EvalMode::Legacy, "UTC", false);
let result = cast_array(array, to_type, &options).unwrap();
assert!(
!result.is_null(0),
"cast {input:?} to {to_type}: Comet still trims {pad:?} where Spark returns \
NULL. If this now returns NULL, the parsers have moved to the trim helpers \
-- delete this test and extend `assert_trim_parity` to the timestamp targets."
);
}
}
}
#[test]
fn test_cast_string_to_timestamp_ntz() {
// Helper to reduce boilerplate
fn parse(s: &str, allow_tz: bool) -> Option<i64> {
timestamp_ntz_parser(s, EvalMode::Legacy, allow_tz, false).unwrap()
}
// Basic: "2020-01-01 12:34:56" -> local micros
// days_from_civil(2020,1,1) = 18262; 18262*86400 = 1577836800
// + 12*3600 + 34*60 + 56 = 45296; total = 1577882096s
assert_eq!(
parse("2020-01-01 12:34:56", true),
Some(1_577_882_096_000_000)
);
assert_eq!(
parse("2020-01-01T12:34:56", true),
Some(1_577_882_096_000_000)
);
// With microseconds
assert_eq!(
parse("2020-01-01 12:34:56.123456", true),
Some(1_577_882_096_123_456)
);
// Date only
assert_eq!(parse("2020-01-01", true), Some(1_577_836_800_000_000));
// Timezone discarded (allow_time_zone=true): same result as without TZ
assert_eq!(
parse("2020-01-01T12:34:56Z", true),
Some(1_577_882_096_000_000)
);
assert_eq!(
parse("2020-01-01T12:34:56+05:30", true),
Some(1_577_882_096_000_000)
);
assert_eq!(
parse("2020-01-01T12:34:56-08:00", true),
Some(1_577_882_096_000_000)
);
// Space-separated offset (e.g. "2021-11-22 10:54:27 +08:00")
assert_eq!(
parse("2021-11-22 10:54:27 +08:00", true),
parse("2021-11-22 10:54:27", true)
);
// Timezone rejected (allow_time_zone=false)
assert_eq!(parse("2020-01-01T12:34:56Z", false), None);
assert_eq!(parse("2020-01-01T12:34:56+05:30", false), None);
// Time-only rejected
assert_eq!(parse("T12:34:56", true), None);
assert_eq!(parse("12:34", true), None);
assert_eq!(parse("T2", true), None);
// Invalid -> None in Legacy
assert_eq!(parse("invalid", true), None);
assert_eq!(parse("", true), None);
// Invalid -> Error in ANSI
assert!(timestamp_ntz_parser("invalid", EvalMode::Ansi, true, false).is_err());
assert!(timestamp_ntz_parser("T12:34", EvalMode::Ansi, true, false).is_err());
// Invalid -> None in Try
assert_eq!(
timestamp_ntz_parser("invalid", EvalMode::Try, true, false).unwrap(),
None
);
// DST gap time works for NTZ (pure arithmetic, no DST)
// days_from_civil(2024,3,10) * 86400 + 2*3600 + 30*60 = 1710037800s
assert_eq!(
parse("2024-03-10 02:30:00", true),
Some(1_710_037_800_000_000)
);
// Invalid leap day -> None
assert_eq!(parse("2023-02-29 00:00:00", true), None);
// Valid leap day
assert!(parse("2020-02-29 00:00:00", true).is_some());
}
#[test]
fn test_cast_string_to_timestamp_ntz_array() {
let array: ArrayRef = Arc::new(StringArray::from(vec![
Some("2020-01-01T12:34:56.123456"),
Some("T2"),
Some("2020-01-01"),
None,
Some("invalid"),
Some("2020-06-15T12:30:00Z"),
]));
let result = cast_string_to_timestamp_ntz(&array, EvalMode::Legacy, true, false).unwrap();
let ts_array = result
.as_any()
.downcast_ref::<PrimitiveArray<TimestampMicrosecondType>>()
.unwrap();
assert_eq!(ts_array.len(), 6);
assert!(!ts_array.is_null(0)); // valid
assert!(ts_array.is_null(1)); // time-only -> null
assert!(!ts_array.is_null(2)); // date-only -> valid
assert!(ts_array.is_null(3)); // null input
assert!(ts_array.is_null(4)); // invalid -> null
assert!(!ts_array.is_null(5)); // TZ discarded -> valid
// TZ discarded: "2020-06-15T12:30:00Z" should give same micros as "2020-06-15T12:30:00"
assert_eq!(
ts_array.value(5),
timestamp_ntz_parser("2020-06-15T12:30:00", EvalMode::Legacy, true, false)
.unwrap()
.unwrap()
);
}
#[test]
fn test_cast_string_to_timestamp_ntz_ansi_error() {
let array: ArrayRef = Arc::new(StringArray::from(vec![Some("invalid")]));
let result = cast_string_to_timestamp_ntz(&array, EvalMode::Ansi, true, false);
assert!(result.is_err());
match result.unwrap_err() {
SparkError::InvalidInputInCastToDatetime { to_type, .. } => {
assert_eq!(to_type, "TIMESTAMP_NTZ");
}
other => panic!("Expected InvalidInputInCastToDatetime, got {other:?}"),
}
}
#[test]
fn test_cast_string_to_timestamp_ntz_ansi_invalid_date() {
// 2023-02-29 is parseable but invalid (not a leap year).
// In ANSI mode this must error, not return NULL.
let result = timestamp_ntz_parser("2023-02-29", EvalMode::Ansi, false, false);
assert!(
result.is_err(),
"ANSI mode should error on invalid date 2023-02-29"
);
match result.unwrap_err() {
SparkError::InvalidInputInCastToDatetime { to_type, .. } => {
assert_eq!(to_type, "TIMESTAMP_NTZ");
}
other => panic!("Expected InvalidInputInCastToDatetime, got {other:?}"),
}
// In Legacy mode, same input should return None (null).
let result = timestamp_ntz_parser("2023-02-29", EvalMode::Legacy, false, false);
assert_eq!(result.unwrap(), None);
}
#[test]
fn test_cast_string_to_timestamp_ntz_out_of_range_year() {
for value in ["294249-01-01", "294249-01-01 00:00:00", "-290310-01-01"] {
let array: ArrayRef = Arc::new(StringArray::from(vec![value]));
match cast_string_to_timestamp_ntz(&array, EvalMode::Ansi, true, false) {
Err(SparkError::InvalidInputInCastToDatetime {
value: actual,
from_type,
to_type,
}) => {
assert_eq!(actual, value);
assert_eq!(from_type, "STRING");
assert_eq!(to_type, "TIMESTAMP_NTZ");
}
other => panic!("Expected ANSI cast error for {value}, got {other:?}"),
}
for mode in [EvalMode::Legacy, EvalMode::Try] {
let result = cast_string_to_timestamp_ntz(&array, mode, true, false).unwrap();
assert!(result.is_null(0), "Expected NULL for {value} in {mode:?}");
}
}
}
#[test]
fn test_cast_dict_string_to_timestamp() -> DataFusionResult<()> {
// prepare input data
let keys = Int32Array::from(vec![0, 1]);
let values: ArrayRef = Arc::new(StringArray::from(vec![
Some("2020-01-01T12:34:56.123456"),
Some("T2"),
]));
let dict_array = Arc::new(DictionaryArray::new(keys, values));
let timezone = "UTC".to_string();
// test casting string dictionary array to timestamp array
let cast_options = SparkCastOptions::new(EvalMode::Legacy, &timezone, false);
let result = cast_array(
dict_array,
&DataType::Timestamp(TimeUnit::Microsecond, Some(timezone.clone().into())),
&cast_options,
)?;
assert_eq!(
*result.data_type(),
DataType::Timestamp(TimeUnit::Microsecond, Some(timezone.into()))
);
assert_eq!(result.len(), 2);
Ok(())
}
#[test]
fn extreme_year_boundary_test() {
let tz = &Tz::from_str("UTC").unwrap();
// Long.MaxValue = 9223372036854775807 μs -> 294247-01-10T04:00:54.775807Z
assert_eq!(
timestamp_parser("294247-01-10T04:00:54.775807Z", EvalMode::Legacy, tz, true).unwrap(),
Some(i64::MAX),
);
// Long.MinValue = -9223372036854775808 μs -> -290308-12-21T19:59:05.224192Z
assert_eq!(
timestamp_parser("-290308-12-21T19:59:05.224192Z", EvalMode::Legacy, tz, true).unwrap(),
Some(i64::MIN),
);
// One beyond Long.MaxValue -> null (overflow)
assert_eq!(
timestamp_parser("294247-01-10T04:00:54.775808Z", EvalMode::Legacy, tz, true).unwrap(),
None,
);
// One before Long.MinValue -> null (overflow)
assert_eq!(
timestamp_parser("-290308-12-21T19:59:05.224191Z", EvalMode::Legacy, tz, true).unwrap(),
None,
);
}
#[test]
fn test_leading_whitespace_t_hm() {
let tz = &Tz::from_str("UTC").unwrap();
// Spark 4.0+ rejects leading whitespace for ALL T-prefixed time-only patterns.
for ws_input in &[
" T2:30",
"\tT2:30",
"\nT2:30",
" T2",
"\tT2",
"\nT2",
"\tT1:2:3 +08:00",
" T1:2:3.4 +08:00",
] {
for mode in [EvalMode::Legacy, EvalMode::Try] {
assert!(
timestamp_parser(ws_input, mode, tz, true)
.unwrap()
.is_none(),
"'{ws_input}' should be null in {mode:?} mode on Spark 4.0+"
);
}
// In ANSI mode the same inputs must raise an error (not silently return null).
assert!(
timestamp_parser(ws_input, EvalMode::Ansi, tz, true).is_err(),
"'{ws_input}' should error in ANSI mode on Spark 4.0+"
);
// Spark 3.x trims all whitespace first, so leading whitespace is valid.
assert!(
timestamp_parser(ws_input, EvalMode::Legacy, tz, false)
.unwrap()
.is_some(),
"'{ws_input}' should be valid in Legacy mode on Spark 3.x"
);
}
// Without leading whitespace, these must be valid on all versions.
for ok_input in &["T2:30", "T2", "T1:2:3 +08:00", "T1:2:3.4 +08:00"] {
assert!(
timestamp_parser(ok_input, EvalMode::Legacy, tz, true)
.unwrap()
.is_some(),
"'{ok_input}' should be valid"
);
}
}
#[test]
fn plus_sign_year_test() {
let tz = &Tz::from_str("UTC").unwrap();
// Spark accepts '+year' prefix on full date-time strings for TIMESTAMP casts.
// "+2020-01-01T12:34:56" -> 2020-01-01T12:34:56 UTC = 1577882096 seconds.
assert_eq!(
timestamp_parser("+2020-01-01T12:34:56", EvalMode::Legacy, tz, true).unwrap(),
Some(1577882096000000),
"+year on full datetime should parse the same as without the + prefix"
);
// But '+' on a time-only string is rejected (Spark returns null).
assert_eq!(
timestamp_parser("+12:12:12", EvalMode::Legacy, tz, true).unwrap(),
None,
"+hour:min:sec must return null"
);
}
#[test]
#[cfg_attr(miri, ignore)] // test takes too long with miri
fn timestamp_parser_test() {
let tz = &Tz::from_str("UTC").unwrap();
// write for all formats
assert_eq!(
timestamp_parser("2020", EvalMode::Legacy, tz, true).unwrap(),
Some(1577836800000000) // this is in milliseconds
);
assert_eq!(
timestamp_parser("2020-01", EvalMode::Legacy, tz, true).unwrap(),
Some(1577836800000000)
);
assert_eq!(
timestamp_parser("2020-01-01", EvalMode::Legacy, tz, true).unwrap(),
Some(1577836800000000)
);
assert_eq!(
timestamp_parser("2020-01-01T12", EvalMode::Legacy, tz, true).unwrap(),
Some(1577880000000000)
);
assert_eq!(
timestamp_parser("2020-01-01T12:34", EvalMode::Legacy, tz, true).unwrap(),
Some(1577882040000000)
);
assert_eq!(
timestamp_parser("2020-01-01T12:34:56", EvalMode::Legacy, tz, true).unwrap(),
Some(1577882096000000)
);
assert_eq!(
timestamp_parser("2020-01-01T12:34:56.123456", EvalMode::Legacy, tz, true).unwrap(),
Some(1577882096123456)
);
assert_eq!(
timestamp_parser("0100", EvalMode::Legacy, tz, true).unwrap(),
Some(-59011459200000000)
);
assert_eq!(
timestamp_parser("0100-01", EvalMode::Legacy, tz, true).unwrap(),
Some(-59011459200000000)
);
assert_eq!(
timestamp_parser("0100-01-01", EvalMode::Legacy, tz, true).unwrap(),
Some(-59011459200000000)
);
assert_eq!(
timestamp_parser("0100-01-01T12", EvalMode::Legacy, tz, true).unwrap(),
Some(-59011416000000000)
);
assert_eq!(
timestamp_parser("0100-01-01T12:34", EvalMode::Legacy, tz, true).unwrap(),
Some(-59011413960000000)
);
assert_eq!(
timestamp_parser("0100-01-01T12:34:56", EvalMode::Legacy, tz, true).unwrap(),
Some(-59011413904000000)
);
assert_eq!(
timestamp_parser("0100-01-01T12:34:56.123456", EvalMode::Legacy, tz, true).unwrap(),
Some(-59011413903876544)
);
assert_eq!(
timestamp_parser("10000", EvalMode::Legacy, tz, true).unwrap(),
Some(253402300800000000)
);
assert_eq!(
timestamp_parser("10000-01", EvalMode::Legacy, tz, true).unwrap(),
Some(253402300800000000)
);
assert_eq!(
timestamp_parser("10000-01-01", EvalMode::Legacy, tz, true).unwrap(),
Some(253402300800000000)
);
assert_eq!(
timestamp_parser("10000-01-01T12", EvalMode::Legacy, tz, true).unwrap(),
Some(253402344000000000)
);
assert_eq!(
timestamp_parser("10000-01-01T12:34", EvalMode::Legacy, tz, true).unwrap(),
Some(253402346040000000)
);
assert_eq!(
timestamp_parser("10000-01-01T12:34:56", EvalMode::Legacy, tz, true).unwrap(),
Some(253402346096000000)
);
assert_eq!(
timestamp_parser("10000-01-01T12:34:56.123456", EvalMode::Legacy, tz, true).unwrap(),
Some(253402346096123456)
);
// Space separator (same values as T separator)
assert_eq!(
timestamp_parser("2020-01-01 12", EvalMode::Legacy, tz, true).unwrap(),
Some(1577880000000000)
);
assert_eq!(
timestamp_parser("2020-01-01 12:34", EvalMode::Legacy, tz, true).unwrap(),
Some(1577882040000000)
);
assert_eq!(
timestamp_parser("2020-01-01 12:34:56", EvalMode::Legacy, tz, true).unwrap(),
Some(1577882096000000)
);
assert_eq!(
timestamp_parser("2020-01-01 12:34:56.123456", EvalMode::Legacy, tz, true).unwrap(),
Some(1577882096123456)
);
// Z suffix (UTC)
assert_eq!(
timestamp_parser("2020-01-01T12:34:56Z", EvalMode::Legacy, tz, true).unwrap(),
Some(1577882096000000)
);
// Positive offset suffix
assert_eq!(
timestamp_parser("2020-01-01T12:34:56+05:30", EvalMode::Legacy, tz, true).unwrap(),
Some(1577862296000000) // 12:34:56 UTC+5:30 = 07:04:56 UTC
);
// T-prefixed time-only with colon
assert!(timestamp_parser("T12:34", EvalMode::Legacy, tz, true)
.unwrap()
.is_some());
assert!(timestamp_parser("T12:34:56", EvalMode::Legacy, tz, true)
.unwrap()
.is_some());
assert!(
timestamp_parser("T12:34:56.123456", EvalMode::Legacy, tz, true)
.unwrap()
.is_some()
);
// Bare time-only (hour:minute without T prefix)
assert!(timestamp_parser("12:34", EvalMode::Legacy, tz, true)
.unwrap()
.is_some());
assert!(timestamp_parser("12:34:56", EvalMode::Legacy, tz, true)
.unwrap()
.is_some());
// Negative year
assert!(timestamp_parser("-0001", EvalMode::Legacy, tz, true)
.unwrap()
.is_some());
assert!(
timestamp_parser("-0001-01-01T12:34:56", EvalMode::Legacy, tz, true)
.unwrap()
.is_some()
);
}
#[test]
#[cfg_attr(miri, ignore)]
fn timestamp_parser_fraction_scaling_test() {
let tz = &Tz::from_str("UTC").unwrap();
// Base: "2020-01-01T12:34:56" = 1577882096000000 µs (confirmed by timestamp_parser_test)
let base = 1577882096000000i64;
// 3-digit fraction: ".123" -> 123_000 µs
assert_eq!(
timestamp_parser("2020-01-01T12:34:56.123", EvalMode::Legacy, tz, true).unwrap(),
Some(base + 123_000)
);
// 1-digit fraction: ".1" -> 100_000 µs
assert_eq!(
timestamp_parser("2020-01-01T12:34:56.1", EvalMode::Legacy, tz, true).unwrap(),
Some(base + 100_000)
);
// 4-digit fraction: ".1000" -> 100_000 µs (trailing zeros not extra precision)
assert_eq!(
timestamp_parser("2020-01-01T12:34:56.1000", EvalMode::Legacy, tz, true).unwrap(),
Some(base + 100_000)
);
// 5-digit fraction: ".12312" -> 123_120 µs
assert_eq!(
timestamp_parser("2020-01-01T12:34:56.12312", EvalMode::Legacy, tz, true).unwrap(),
Some(base + 123_120)
);
// 6-digit fraction (exact): unchanged
assert_eq!(
timestamp_parser("2020-01-01T12:34:56.123456", EvalMode::Legacy, tz, true).unwrap(),
Some(base + 123_456)
);
// >6 digits: truncated to 6
assert_eq!(
timestamp_parser("2020-01-01T12:34:56.123456789", EvalMode::Legacy, tz, true).unwrap(),
Some(base + 123_456)
);
// Fraction after Z-stripped offset
assert_eq!(
timestamp_parser("2020-01-01T12:34:56.123Z", EvalMode::Legacy, tz, true).unwrap(),
Some(base + 123_000)
);
}
#[test]
#[cfg_attr(miri, ignore)]
fn timestamp_parser_tz_offset_formats_test() {
let tz = &Tz::from_str("UTC").unwrap();
// All of these represent 2020-01-01T12:34:56 UTC = 1577882096000000 µs.
let utc = 1577882096000000i64;
// +05:30 offset -> UTC = 12:34:56 − 5h30m = 07:04:56 UTC = 1577862296000000 µs
let plus530 = 1577862296000000i64;
// +/-HHMM (no colon)
assert_eq!(
timestamp_parser("2020-01-01T12:34:56+0000", EvalMode::Legacy, tz, true).unwrap(),
Some(utc)
);
assert_eq!(
timestamp_parser("2020-01-01T12:34:56+0530", EvalMode::Legacy, tz, true).unwrap(),
Some(plus530)
);
// +/-H:MM (single-digit hour)
assert_eq!(
timestamp_parser("2020-01-01T12:34:56+5:30", EvalMode::Legacy, tz, true).unwrap(),
Some(plus530)
);
assert_eq!(
timestamp_parser("2020-01-01T12:34:56+0:00", EvalMode::Legacy, tz, true).unwrap(),
Some(utc)
);
// +/-H:M (single-digit both)
assert_eq!(
timestamp_parser("2020-01-01T12:34:56+5:3", EvalMode::Legacy, tz, true).unwrap(),
Some(1577863916000000) // 12:34:56 − 5h3m = 07:31:56 UTC = 1577836800+27116
);
// bare UTC / " UTC"
assert_eq!(
timestamp_parser("2020-01-01T12:34:56UTC", EvalMode::Legacy, tz, true).unwrap(),
Some(utc)
);
assert_eq!(
timestamp_parser("2020-01-01T12:34:56 UTC", EvalMode::Legacy, tz, true).unwrap(),
Some(utc)
);
// UTC+offset
assert_eq!(
timestamp_parser("2020-01-01T12:34:56 UTC+5:30", EvalMode::Legacy, tz, true).unwrap(),
Some(plus530)
);
// UTC+0 (single-digit zero)
assert_eq!(
timestamp_parser("2020-01-01T12:34:56 UTC+0", EvalMode::Legacy, tz, true).unwrap(),
Some(utc)
);
// GMT+/-HH:MM (no space)
assert_eq!(
timestamp_parser("2020-01-01T12:34:56GMT+00:00", EvalMode::Legacy, tz, true).unwrap(),
Some(utc)
);
assert_eq!(
timestamp_parser("2020-01-01T12:34:56GMT+05:30", EvalMode::Legacy, tz, true).unwrap(),
Some(plus530)
);
// " GMT+/-..." (space-prefixed)
assert_eq!(
timestamp_parser("2020-01-01T12:34:56 GMT+05:30", EvalMode::Legacy, tz, true).unwrap(),
Some(plus530)
);
// " GMT+/-HHMM" (space + GMT + no colon)
assert_eq!(
timestamp_parser("2020-01-01T12:34:56 GMT+0530", EvalMode::Legacy, tz, true).unwrap(),
Some(plus530)
);
// " UT+/-HH:MM"
assert_eq!(
timestamp_parser("2020-01-01T12:34:56 UT+05:30", EvalMode::Legacy, tz, true).unwrap(),
Some(plus530)
);
// Bare "UT" (no leading space) — Spark accepts "UT" as a UTC alias.
assert_eq!(
timestamp_parser("2020-01-01T12:34:56UT", EvalMode::Legacy, tz, true).unwrap(),
Some(utc)
);
// Java SHORT_IDS: EST (-05:00), MST (-07:00), HST (-10:00)
// 2020-01-01T12:34:56 EST = 2020-01-01T17:34:56 UTC = 1577896496 s
let est_utc = utc + 5 * 3600 * 1_000_000;
assert_eq!(
timestamp_parser("2020-01-01T12:34:56 EST", EvalMode::Legacy, tz, true).unwrap(),
Some(est_utc)
);
assert_eq!(
timestamp_parser("2020-01-01T12:34:56EST", EvalMode::Legacy, tz, true).unwrap(),
Some(est_utc)
);
// 2020-01-01T12:34:56 MST = 2020-01-01T19:34:56 UTC
let mst_utc = utc + 7 * 3600 * 1_000_000;
assert_eq!(
timestamp_parser("2020-01-01T12:34:56 MST", EvalMode::Legacy, tz, true).unwrap(),
Some(mst_utc)
);
// 2020-01-01T12:34:56 HST = 2020-01-01T22:34:56 UTC
let hst_utc = utc + 10 * 3600 * 1_000_000;
assert_eq!(
timestamp_parser("2020-01-01T12:34:56 HST", EvalMode::Legacy, tz, true).unwrap(),
Some(hst_utc)
);
// Named IANA zone " Europe/Moscow" (UTC+3 in winter 2020)
// 2020-01-01T12:34:56 Europe/Moscow = 2020-01-01T09:34:56 UTC = 1577871296000000 µs
assert_eq!(
timestamp_parser(
"2020-01-01T12:34:56 Europe/Moscow",
EvalMode::Legacy,
tz,
true
)
.unwrap(),
Some(1577871296000000)
);
// Plain date strings must NOT be affected by the offset-extraction logic.
assert_eq!(
timestamp_parser("2020-01-01", EvalMode::Legacy, tz, true).unwrap(),
Some(1577836800000000)
);
// Invalid offset formats -> null
assert_eq!(
timestamp_parser("2020-01-01T12:34:56-8:", EvalMode::Legacy, tz, true).unwrap(),
None
);
assert_eq!(
timestamp_parser("2020-01-01T12:34:56-20:0", EvalMode::Legacy, tz, true).unwrap(),
None // h=20 > 18 invalid
);
// Positive year-sign prefix is accepted for timestamps (see plus_sign_year_test)
assert_eq!(
timestamp_parser("+2020-01-01T12:34:56", EvalMode::Legacy, tz, true).unwrap(),
Some(1577882096000000)
);
}
#[test]
#[cfg_attr(miri, ignore)]
fn timestamp_parser_dst_test() {
// DST spring-forward: America/New_York springs forward 2020-03-08 02:00 -> 03:00.
// 02:30 does not exist; Spark advances to 03:30 EDT (UTC-4) = 07:30 UTC.
// 2020-03-08T07:30:00Z = 1577836800 + 67*86400 + 27000 = 1583652600 seconds.
let ny_tz = &Tz::from_str("America/New_York").unwrap();
assert_eq!(
timestamp_parser("2020-03-08 02:30:00", EvalMode::Legacy, ny_tz, true).unwrap(),
Some(1583652600000000)
);
// Just before gap: 01:59:59 EST (UTC-5) = 06:59:59 UTC = 1583650799 seconds.
assert_eq!(
timestamp_parser("2020-03-08 01:59:59", EvalMode::Legacy, ny_tz, true).unwrap(),
Some(1583650799000000)
);
// Just after gap: 03:00:00 EDT (UTC-4) = 07:00:00 UTC = 1583650800 seconds.
assert_eq!(
timestamp_parser("2020-03-08 03:00:00", EvalMode::Legacy, ny_tz, true).unwrap(),
Some(1583650800000000)
);
// DST fall-back: 2020-11-01 02:00 EDT -> 01:00 EST. Ambiguous: [01:00, 02:00).
// Spark picks the earlier UTC instant (pre-transition = EDT = UTC-4).
// 01:30 EDT (UTC-4) = 05:30 UTC.
// 2020-11-01 = 2020-01-01 + 305 days = 1577836800 + 305*86400 = 1604188800 seconds.
assert_eq!(
timestamp_parser("2020-11-01 01:30:00", EvalMode::Legacy, ny_tz, true).unwrap(),
Some(1604208600000000) // 1604188800 + 5*3600 + 30*60 = 1604208600
);
}
// 2020-01-01T00:00:00Z, 2020-01-01T12:34:56Z and the same wall clock at +05:30, in micros.
const JAN1_2020: i64 = 1577836800000000;
const JAN1_2020_123456: i64 = 1577882096000000;
const JAN1_2020_123456_PLUS_0530: i64 = 1577862296000000;
/// Inputs Spark's `parseTimestampString` accepts that the fixed 2-digit shapes rejected:
/// 1-2 digit month/day/hour/minute/second, an empty fraction (also before a zone), and
/// 6-digit years (issue #5674). Values are UTC micros, identical for TIMESTAMP_NTZ.
const SPARK_SEGMENT_RULE_VALID: &[(&str, i64)] = &[
("2020-10-1", 1_601_510_400_000_000),
("2020-12-1", 1_606_780_800_000_000),
("2020-1", JAN1_2020),
("2020-1-1", JAN1_2020),
("2020-1-1T1", JAN1_2020 + 3600 * 1_000_000),
("2020-1-1 1:2", JAN1_2020 + 3720 * 1_000_000),
("2020-01-01 12:34:5", JAN1_2020 + 45245 * 1_000_000),
("2020-1-1T1:2:3.4", JAN1_2020 + 3723 * 1_000_000 + 400_000),
("2020-01-01 12:34:56.", JAN1_2020_123456),
("002020-01-01 00:00:00", JAN1_2020),
];
/// Inputs Spark rejects: non-ASCII segment digits, a zone suffix anywhere but after the
/// seconds segment, more than six year digits, and more than two digits in any other segment.
const SPARK_SEGMENT_RULE_INVALID: &[&str] = &[
"2020-01-01 12:34:56.1٢٢٢",
"T1:2:3.1٢٢٢",
"2020-1-1TÙ¢",
"2020-1-1T1:2:3.Ù¢",
"Ù¢020-1-1",
"2020-Ù¢",
"2020-01-Ù¢",
"2020-1\u{967}",
"2020-\u{967}1",
"2020-01-1\u{967}",
"2020-01-\u{967}1",
"2020-01-01 12:34:56 +08:000",
"2020-01-01 12:34:56 +008:00",
"2020-01-01 12:34:56 +18:01",
"2020-01-01 12:34:56 -18:01",
"2020-01-01 12:34:56+08:000",
"2020-01-01 12:34:56+008:00",
"2020-01-01 12:34:56+18:01",
"2020-01-01 12:34:56 UTC+08:000",
"2020-01-01 12:34:56 GMT+008:00",
"2020-01-01 12:34:56 UT+18:01",
"2020-01-01T1:Ù¢",
"2020-01-01T1:2:Ù£",
"2020-1-1T1:2:3.Ù¢Z",
"TÙ¢",
"T1:Ù¢",
"T1:2:Ù£",
"T1:2:3.Ù¢",
"1:Ù¢",
"1:2:Ù£",
"1:2:3.Ù¢",
"2020Z",
"2020-10-01Z",
"2020-01-01+05:30",
"2020-01-01-08:00",
"2020-10-01 UTC",
"2020-01-01T12Z",
"2020-01-01 12 UTC",
"2020-01-01T12:34Z",
"2020-01-01 12:34 UTC",
"2020-01-01T12:34:Z",
"0002020-01-01",
"0002020-01-01 00:00:00",
"-0002020-01-01",
"2020-001-01",
"2020-01-001",
"2020-01-01T123",
"2020-01-01T12:345",
"2020-01-01T12:34:567",
];
#[test]
#[cfg_attr(miri, ignore)]
fn timestamp_zone_whitespace_matches_java_trim() {
let tz = &Tz::from_str("UTC").unwrap();
for whitespace in [
" ", "\t", "\n", "\u{1}", "\u{b}", "\u{c}", "\u{7f}", "\u{a0}", "\u{2009}", "\u{3000}",
] {
let valid = whitespace.chars().all(|c| c <= '\u{20}');
for (suffix, offset) in [("+08:00", 28_800_000_000), ("UTC", 0), ("Z", 0)] {
for (fraction, micros) in [("", 0), (".123", 123_000)] {
let input = format!("2020-01-01 12:34:56{fraction}{whitespace}{suffix}");
for mode in [EvalMode::Legacy, EvalMode::Try, EvalMode::Ansi] {
for spark4 in [false, true] {
for (result, expected) in [
(
timestamp_parser(&input, mode, tz, spark4),
JAN1_2020_123456 + micros - offset,
),
(
timestamp_ntz_parser(&input, mode, true, spark4),
JAN1_2020_123456 + micros,
),
] {
if valid {
assert_eq!(
result.unwrap(),
Some(expected),
"{input:?}, {mode:?}"
);
} else if mode == EvalMode::Ansi {
assert!(
matches!(
result,
Err(SparkError::InvalidInputInCastToDatetime { .. })
),
"{input:?}"
);
} else {
assert_eq!(result.unwrap(), None, "{input:?}, {mode:?}");
}
}
let no_zone = timestamp_ntz_parser(&input, mode, false, spark4);
if mode == EvalMode::Ansi {
assert!(no_zone.is_err(), "{input:?}");
} else {
assert_eq!(no_zone.unwrap(), None, "{input:?}");
}
}
}
}
}
}
}
#[test]
fn timestamp_numeric_offset_validation() {
for (input, seconds) in [
("", 0),
("+0", 0),
("-00", 0),
("+8", 28_800),
("+08", 28_800),
("+0800", 28_800),
("+8:0", 28_800),
("+08:0", 28_800),
("+8:00", 28_800),
("+17:59", 64_740),
("+18:00", 64_800),
("-18:00", -64_800),
("+1800", 64_800),
("-1800", -64_800),
] {
assert_eq!(parse_sign_offset(input), Some(seconds), "{input:?}");
}
for input in [
"+08:000",
"+008:00",
"+18:01",
"-18:01",
"+18:59",
"+19",
"-1900",
"+8:",
"+1:+1",
"-1:-1",
"+1\u{967}",
"+\u{967}1",
"+1Ù¢",
"++1",
] {
assert_eq!(parse_sign_offset(input), None, "{input:?}");
}
}
#[test]
#[cfg_attr(miri, ignore)]
fn timestamp_parser_spark_segment_rules_test() {
let tz = &Tz::from_str("UTC").unwrap();
// Exercise the decoders without their regex gates: malformed UTF-8 boundaries
// must not panic even if the accepted patterns change later.
assert!(
parse_to_timestamp_info("2020-01-01 12:34:56.1٢٢٢", "microsecond")
.unwrap()
.is_none()
);
assert_eq!(
parse_str_to_time_only_timestamp("T1:2:3.1٢٢٢", tz).unwrap(),
None
);
for mode in [EvalMode::Legacy, EvalMode::Try, EvalMode::Ansi] {
assert_eq!(
timestamp_parser("2021-11-22 10:54:27 +08:00", mode, tz, true).unwrap(),
Some(1_637_549_667_000_000)
);
assert_eq!(
timestamp_ntz_parser("2021-11-22 10:54:27 +08:00", mode, true, true).unwrap(),
Some(1_637_578_467_000_000)
);
}
for &(input, expected) in SPARK_SEGMENT_RULE_VALID {
for eval_mode in [EvalMode::Legacy, EvalMode::Try, EvalMode::Ansi] {
assert_eq!(
timestamp_parser(input, eval_mode, tz, true).unwrap(),
Some(expected),
"{input:?} in {eval_mode:?}"
);
}
}
// An empty fraction may still be followed by a zone.
assert_eq!(
timestamp_parser("2020-01-01 12:34:56.Z", EvalMode::Legacy, tz, true).unwrap(),
Some(JAN1_2020_123456)
);
assert_eq!(
timestamp_parser("2020-01-01 12:34:56.+05:30", EvalMode::Legacy, tz, true).unwrap(),
Some(JAN1_2020_123456_PLUS_0530)
);
// Time-only shapes may carry a zone after their seconds segment but not before it.
for input in ["T12:34:56Z", "12:34:56+05:30", "T1:2:3.Z"] {
assert!(
timestamp_parser(input, EvalMode::Ansi, tz, true)
.unwrap()
.is_some(),
"{input:?}"
);
}
for input in
SPARK_SEGMENT_RULE_INVALID
.iter()
.copied()
.chain(["T12Z", "12:34Z", "T12:34 UTC"])
{
for eval_mode in [EvalMode::Legacy, EvalMode::Try] {
assert_eq!(
timestamp_parser(input, eval_mode, tz, true).unwrap(),
None,
"{input:?} in {eval_mode:?}"
);
}
assert!(
timestamp_parser(input, EvalMode::Ansi, tz, true).is_err(),
"{input:?} in Ansi"
);
}
// Shapes that were already accepted keep their exact values.
let la_offset = 8 * 3600 * 1_000_000; // America/Los_Angeles is UTC-8 in January
for (input, expected) in [
("2020-01-01", JAN1_2020),
("2020-01-01 12:34:56", JAN1_2020_123456),
("2020-01-01T12:34:56.123456", JAN1_2020_123456 + 123456),
("2020-01-01T12:34:56Z", JAN1_2020_123456),
("2020-01-01T12:34:56.123Z", JAN1_2020_123456 + 123000),
("2020-01-01T12:34:56+05:30", JAN1_2020_123456_PLUS_0530),
("2020-01-01T12:34:56 UTC", JAN1_2020_123456),
("2020-01-01T12:34:56 UTC+5:30", JAN1_2020_123456_PLUS_0530),
(
"2020-01-01T12:34:56 America/Los_Angeles",
JAN1_2020_123456 + la_offset,
),
("-0001-01-01T12:34:56", -62198709904000000),
] {
for eval_mode in [EvalMode::Legacy, EvalMode::Try, EvalMode::Ansi] {
assert_eq!(
timestamp_parser(input, eval_mode, tz, true).unwrap(),
Some(expected),
"{input:?} in {eval_mode:?}"
);
}
}
// `date_parser` ports `stringToDate`, whose `maxDigitsYear` is 7, so a date cast keeps
// accepting the 7-digit year that a timestamp cast rejects.
for eval_mode in [EvalMode::Legacy, EvalMode::Try, EvalMode::Ansi] {
assert_eq!(
date_parser("0002020-01-01", eval_mode).unwrap(),
Some(18262)
);
}
}
#[test]
#[cfg_attr(miri, ignore)]
fn timestamp_ntz_parser_spark_segment_rules_test() {
for allow_time_zone in [true, false] {
for &(input, expected) in SPARK_SEGMENT_RULE_VALID {
for eval_mode in [EvalMode::Legacy, EvalMode::Try, EvalMode::Ansi] {
assert_eq!(
timestamp_ntz_parser(input, eval_mode, allow_time_zone, false).unwrap(),
Some(expected),
"{input:?} in {eval_mode:?}, allow_time_zone={allow_time_zone}"
);
}
}
for &input in SPARK_SEGMENT_RULE_INVALID {
for eval_mode in [EvalMode::Legacy, EvalMode::Try] {
assert_eq!(
timestamp_ntz_parser(input, eval_mode, allow_time_zone, false).unwrap(),
None,
"{input:?} in {eval_mode:?}, allow_time_zone={allow_time_zone}"
);
}
assert!(
timestamp_ntz_parser(input, EvalMode::Ansi, allow_time_zone, false).is_err(),
"{input:?} in Ansi, allow_time_zone={allow_time_zone}"
);
}
}
// A zone after an empty fraction is discarded when allowed and rejected otherwise.
assert_eq!(
timestamp_ntz_parser("2020-01-01 12:34:56.Z", EvalMode::Legacy, true, false).unwrap(),
Some(JAN1_2020_123456)
);
assert_eq!(
timestamp_ntz_parser("2020-01-01 12:34:56.Z", EvalMode::Legacy, false, false).unwrap(),
None
);
assert!(
timestamp_ntz_parser("2020-01-01 12:34:56.Z", EvalMode::Ansi, false, false).is_err()
);
}
#[test]
#[cfg_attr(miri, ignore)]
fn test_cast_string_to_timestamp_spark_segment_rules_array() {
// The reproducer from issue #5674, through the batch entry points.
let inputs = vec![
Some("2020-1-1"),
Some("2020-01-01 12:34:5"),
Some("2020-01-01 12:34:56."),
Some("2020-10-01Z"),
Some("0002020-01-01 00:00:00"),
];
let expected = [
Some(JAN1_2020),
Some(JAN1_2020 + 45245 * 1_000_000),
Some(JAN1_2020_123456),
None,
None,
];
let array: ArrayRef = Arc::new(StringArray::from(inputs));
let to_type = DataType::Timestamp(TimeUnit::Microsecond, Some("UTC".into()));
let tz_result =
cast_string_to_timestamp(&array, &to_type, EvalMode::Legacy, "UTC", true).unwrap();
let ntz_result =
cast_string_to_timestamp_ntz(&array, EvalMode::Legacy, true, false).unwrap();
for result in [&tz_result, &ntz_result] {
let result = result
.as_any()
.downcast_ref::<PrimitiveArray<TimestampMicrosecondType>>()
.unwrap();
let actual: Vec<Option<i64>> = result.iter().collect();
assert_eq!(actual, expected);
}
// Under ANSI the first malformed row fails the batch and names the raw input.
let tz_err =
cast_string_to_timestamp(&array, &to_type, EvalMode::Ansi, "UTC", true).unwrap_err();
let ntz_err =
cast_string_to_timestamp_ntz(&array, EvalMode::Ansi, true, false).unwrap_err();
for (err, expected_type) in [(tz_err, "TIMESTAMP"), (ntz_err, "TIMESTAMP_NTZ")] {
match err {
SparkError::InvalidInputInCastToDatetime {
value,
from_type,
to_type,
} => {
assert_eq!(value, "2020-10-01Z");
assert_eq!(from_type, "STRING");
assert_eq!(to_type, expected_type);
}
other => panic!("Expected InvalidInputInCastToDatetime, got {other:?}"),
}
}
}
/// Asserts every date parses to null in legacy and try mode. When `expect_ansi_error` is set,
/// ANSI mode must raise CAST_INVALID_INPUT; otherwise ANSI mode must also return null.
fn assert_dates(dates: &[&str], expect_ansi_error: bool) {
for &date in dates {
for eval_mode in [EvalMode::Legacy, EvalMode::Try, EvalMode::Ansi] {
if expect_ansi_error && eval_mode == EvalMode::Ansi {
assert!(date_parser(date, eval_mode).is_err(), "{date}");
} else {
assert_eq!(date_parser(date, eval_mode).unwrap(), None, "{date}");
}
}
}
}
/// Malformed input: null in legacy and try mode, CAST_INVALID_INPUT in ANSI mode.
fn assert_null_or_ansi_error(dates: &[&str]) {
assert_dates(dates, true);
}
/// Input Spark parses successfully but Comet cannot represent: null in every eval mode.
fn assert_null_in_all_modes(dates: &[&str]) {
assert_dates(dates, false);
}
#[test]
fn date_parser_test() {
for date in &[
"2020",
"2020-01",
"2020-01-01",
"+2020-01-01", // Spark accepts '+' year prefix on dates
"02020-01-01",
"002020-01-01",
"0002020-01-01",
"2020-1-1",
"2020-01-01 ",
"2020-01-01T",
] {
for eval_mode in &[EvalMode::Legacy, EvalMode::Ansi, EvalMode::Try] {
assert_eq!(date_parser(date, *eval_mode).unwrap(), Some(18262));
}
}
//dates in invalid formats
assert_null_or_ansi_error(&[
"abc",
"",
"not_a_date",
"3/",
"3/12",
"3/12/2020",
"3/12/2002 T",
"202",
"2020-010-01",
"2020-10-010",
"2020-10-010T",
"--262143-12-31",
"--262143-12-31 ",
]);
for date in &["-3638-5"] {
for eval_mode in &[EvalMode::Legacy, EvalMode::Try, EvalMode::Ansi] {
assert_eq!(date_parser(date, *eval_mode).unwrap(), Some(-2048160));
}
}
//Naive Date only supports years 262142 AD to 262143 BC. Spark parses these fine, so
//they are a Comet limitation rather than malformed input and stay null in ANSI mode.
assert_null_in_all_modes(&[
"-262144-1-1",
"262143-01-1",
"262143-1-1",
"262143-01-1 ",
"262143-01-01T ",
"262143-1-01T 1234",
"-0973250",
]);
//years whose epoch day overflows i32 are rejected by Spark too (localDateToDays uses
//Math.toIntExact), so ANSI mode must raise rather than return null
assert_null_or_ansi_error(&["9999999-01-01", "-9999999-01-01"]);
// Canonical `yyyy-mm-dd` shape with invalid calendar dates exercises the fast path.
// Spark's LocalDate.of rejects these, so ANSI mode must raise (issue #5012).
assert_null_or_ansi_error(&[
"2020-02-30",
"2021-02-29",
"2020-13-01",
"2020-00-15",
"2020-04-31",
"2020-01-00",
]);
// Same invalid calendar dates in non-canonical shapes take the general parser path.
assert_null_or_ansi_error(&["2020-2-30", "2020-13-1", "2020-4-31 ", "2020-02-30T"]);
// Valid leap day flows through the fast path.
assert_eq!(
date_parser("2020-02-29", EvalMode::Legacy).unwrap(),
Some(18321)
);
}
#[test]
fn test_cast_string_to_date() {
let array: ArrayRef = Arc::new(StringArray::from(vec![
Some("2020"),
Some("2020-01"),
Some("2020-01-01"),
Some("2020-01-01T"),
]));
let result = cast_string_to_date(&array, &DataType::Date32, EvalMode::Legacy).unwrap();
let date32_array = result
.as_any()
.downcast_ref::<arrow::array::Date32Array>()
.unwrap();
assert_eq!(date32_array.len(), 4);
date32_array
.iter()
.for_each(|v| assert_eq!(v.unwrap(), 18262));
}
#[test]
fn test_cast_string_array_with_valid_dates() {
let array_with_invalid_date: ArrayRef = Arc::new(StringArray::from(vec![
Some("-262143-12-31"),
Some("\n -262143-12-31 "),
Some("-262143-12-31T \t\n"),
Some("\n\t-262143-12-31T\r"),
Some("-262143-12-31T 123123123"),
Some("\r\n-262143-12-31T \r123123123"),
Some("\n -262143-12-31T \n\t"),
]));
for eval_mode in &[EvalMode::Legacy, EvalMode::Try, EvalMode::Ansi] {
let result =
cast_string_to_date(&array_with_invalid_date, &DataType::Date32, *eval_mode)
.unwrap();
let date32_array = result
.as_any()
.downcast_ref::<arrow::array::Date32Array>()
.unwrap();
assert_eq!(result.len(), 7);
date32_array
.iter()
.for_each(|v| assert_eq!(v.unwrap(), -96464928));
}
}
#[test]
fn test_cast_string_array_with_invalid_dates() {
let array_with_invalid_date: ArrayRef = Arc::new(StringArray::from(vec![
Some("2020"),
Some("2020-01"),
Some("2020-01-01"),
//4 invalid dates
Some("2020-010-01T"),
Some("202"),
Some(" 202 "),
Some("\n 2020-\r8 "),
Some("2020-01-01T"),
// Overflows i32
Some("-4607172990231812908"),
]));
for eval_mode in &[EvalMode::Legacy, EvalMode::Try] {
let result =
cast_string_to_date(&array_with_invalid_date, &DataType::Date32, *eval_mode)
.unwrap();
let date32_array = result
.as_any()
.downcast_ref::<arrow::array::Date32Array>()
.unwrap();
assert_eq!(
date32_array.iter().collect::<Vec<_>>(),
vec![
Some(18262),
Some(18262),
Some(18262),
None,
None,
None,
None,
Some(18262),
None
]
);
}
let result =
cast_string_to_date(&array_with_invalid_date, &DataType::Date32, EvalMode::Ansi);
match result {
Err(e) => assert!(
e.to_string().contains(
"[CAST_INVALID_INPUT] The value '2020-010-01T' of the type \"STRING\" cannot be cast to \"DATE\" because it is malformed")
),
_ => panic!("Expected error"),
}
}
#[test]
fn test_cast_string_as_i8() {
// basic
assert_eq!(
cast_string_to_i8("127", EvalMode::Legacy).unwrap(),
Some(127_i8)
);
assert_eq!(cast_string_to_i8("128", EvalMode::Legacy).unwrap(), None);
assert!(cast_string_to_i8("128", EvalMode::Ansi).is_err());
// decimals
assert_eq!(
cast_string_to_i8("0.2", EvalMode::Legacy).unwrap(),
Some(0_i8)
);
assert_eq!(
cast_string_to_i8(".", EvalMode::Legacy).unwrap(),
Some(0_i8)
);
// TRY should always return null for decimals
assert_eq!(cast_string_to_i8("0.2", EvalMode::Try).unwrap(), None);
assert_eq!(cast_string_to_i8(".", EvalMode::Try).unwrap(), None);
// ANSI mode should throw error on decimal
assert!(cast_string_to_i8("0.2", EvalMode::Ansi).is_err());
assert!(cast_string_to_i8(".", EvalMode::Ansi).is_err());
}
}