diff --git a/datafusion/functions-nested/src/position.rs b/datafusion/functions-nested/src/position.rs index f677cfb4979d3..9ef8142fa48dd 100644 --- a/datafusion/functions-nested/src/position.rs +++ b/datafusion/functions-nested/src/position.rs @@ -40,7 +40,9 @@ use arrow::array::{ use datafusion_common::cast::{ as_generic_list_array, as_int64_array, as_large_list_array, as_list_array, }; -use datafusion_common::{Result, exec_err, utils::take_function_args}; +use datafusion_common::{ + Result, exec_datafusion_err, exec_err, utils::take_function_args, +}; use itertools::Itertools; use crate::utils::{compare_element_to_list, make_scalar_function}; @@ -198,6 +200,16 @@ fn array_position_inner(args: &[ArrayRef]) -> Result { } } +fn resolve_zero_based_start_from(start_from: i64) -> Result { + start_from.checked_sub(1).ok_or_else(|| { + exec_datafusion_err!( + "start_from out of bounds: {start_from}, expected {} to {}", + i64::MIN + 1, + i64::MAX + ) + }) +} + /// Resolves the optional `start_from` argument into a `Vec` of /// 0-indexed starting positions. fn resolve_start_from( @@ -207,14 +219,16 @@ fn resolve_start_from( match third_arg { None => Ok(vec![0i64; num_rows]), Some(ColumnarValue::Scalar(ScalarValue::Int64(Some(v)))) => { - Ok(vec![v - 1; num_rows]) + Ok(vec![resolve_zero_based_start_from(*v)?; num_rows]) } Some(ColumnarValue::Scalar(s)) => { exec_err!("array_position expected Int64 for start_from, got {s}") } - Some(ColumnarValue::Array(a)) => { - Ok(as_int64_array(a)?.values().iter().map(|&x| x - 1).collect()) - } + Some(ColumnarValue::Array(a)) => as_int64_array(a)? + .values() + .iter() + .map(|&x| resolve_zero_based_start_from(x)) + .collect(), } } @@ -309,8 +323,8 @@ fn general_position_dispatch(args: &[ArrayRef]) -> Result>() + .map(|&x| resolve_zero_based_start_from(x)) + .collect::>>()? } else { vec![0; haystack.len()] }; diff --git a/datafusion/sqllogictest/test_files/array/array_position.slt b/datafusion/sqllogictest/test_files/array/array_position.slt index e3dd830dfb77a..8fe1826619431 100644 --- a/datafusion/sqllogictest/test_files/array/array_position.slt +++ b/datafusion/sqllogictest/test_files/array/array_position.slt @@ -282,6 +282,17 @@ select array_position([1, 2, 3], 3, 4), array_position([1], 1, 2); ---- NULL NULL +query error start_from out of bounds: -9223372036854775808, expected -9223372036854775807 to 9223372036854775807 +select array_position([1], 1, -9223372036854775808); + +query error start_from out of bounds: -9223372036854775808, expected -9223372036854775807 to 9223372036854775807 +select array_position([1], 1, start_from) +from (values (-9223372036854775808)) as t(start_from); + +query error start_from out of bounds: -9223372036854775808, expected -9223372036854775807 to 9223372036854775807 +select array_position([1], needle, start_from) +from (values (1, -9223372036854775808)) as t(needle, start_from); + # array_position with empty array in various contexts query II select array_position(arrow_cast(make_array(), 'List(Int64)'), 1), array_position(arrow_cast(make_array(), 'LargeList(Int64)'), 1);