diff --git a/src/common/function/src/aggrs/geo/encoding.rs b/src/common/function/src/aggrs/geo/encoding.rs index 35a59d1ecfb..67dfd493a77 100644 --- a/src/common/function/src/aggrs/geo/encoding.rs +++ b/src/common/function/src/aggrs/geo/encoding.rs @@ -221,18 +221,24 @@ impl DfAccumulator for JsonEncodePathAccumulator { ))); } - for state in states { - let state = as_struct_array(state)?; - let lat_list = as_list_array(state.column(0))?.value(0); - let lat_array = as_primitive_array::(&lat_list)?; - let lng_list = as_list_array(state.column(1))?.value(0); - let lng_array = as_primitive_array::(&lng_list)?; - let ts_list = as_list_array(state.column(2))?.value(0); - let ts_array = as_primitive_array::(&ts_list)?; + let state = as_struct_array(&states[0])?; + let lat_lists = as_list_array(state.column(0))?; + let lng_lists = as_list_array(state.column(1))?; + let ts_lists = as_list_array(state.column(2))?; + for row in 0..state.len() { + if state.is_null(row) { + continue; + } + let lat_list = lat_lists.value(row); + let lng_list = lng_lists.value(row); + let ts_list = ts_lists.value(row); - self.lat.extend(lat_array); - self.lng.extend(lng_array); - self.timestamp.extend(ts_array); + self.lat + .extend(as_primitive_array::(&lat_list)?); + self.lng + .extend(as_primitive_array::(&lng_list)?); + self.timestamp + .extend(as_primitive_array::(&ts_list)?); } Ok(()) @@ -331,9 +337,9 @@ mod tests { _ => panic!("Expected Struct scalar value"), }; - // Merge state arrays - merged.merge_batch(&[state_array1]).unwrap(); - merged.merge_batch(&[state_array2]).unwrap(); + // States from different partitions arrive as rows of one batch + let states = compute::concat(&[state_array1.as_ref(), state_array2.as_ref()]).unwrap(); + merged.merge_batch(&[states]).unwrap(); // Evaluate merged result let result = merged.evaluate().unwrap(); diff --git a/src/common/function/src/aggrs/geo/geo_path.rs b/src/common/function/src/aggrs/geo/geo_path.rs index b09745eaeb7..c4e366c83b0 100644 --- a/src/common/function/src/aggrs/geo/geo_path.rs +++ b/src/common/function/src/aggrs/geo/geo_path.rs @@ -250,18 +250,24 @@ impl DfAccumulator for GeoPathAccumulator { ))); } - for state in states { - let state = as_struct_array(state)?; - let lat_list = as_list_array(state.column(0))?.value(0); - let lat_array = as_primitive_array::(&lat_list)?; - let lng_list = as_list_array(state.column(1))?.value(0); - let lng_array = as_primitive_array::(&lng_list)?; - let ts_list = as_list_array(state.column(2))?.value(0); - let ts_array = as_primitive_array::(&ts_list)?; + let state = as_struct_array(&states[0])?; + let lat_lists = as_list_array(state.column(0))?; + let lng_lists = as_list_array(state.column(1))?; + let ts_lists = as_list_array(state.column(2))?; + for row in 0..state.len() { + if state.is_null(row) { + continue; + } + let lat_list = lat_lists.value(row); + let lng_list = lng_lists.value(row); + let ts_list = ts_lists.value(row); - self.lat.extend(lat_array); - self.lng.extend(lng_array); - self.timestamp.extend(ts_array); + self.lat + .extend(as_primitive_array::(&lat_list)?); + self.lng + .extend(as_primitive_array::(&lng_list)?); + self.timestamp + .extend(as_primitive_array::(&ts_list)?); } Ok(()) @@ -403,9 +409,9 @@ mod tests { _ => panic!("Expected Struct scalar value"), }; - // Merge state arrays - merged.merge_batch(&[state_array1]).unwrap(); - merged.merge_batch(&[state_array2]).unwrap(); + // States from different partitions arrive as rows of one batch + let states = compute::concat(&[state_array1.as_ref(), state_array2.as_ref()]).unwrap(); + merged.merge_batch(&[states]).unwrap(); // Evaluate merged result let result = merged.evaluate().unwrap(); diff --git a/tests/cases/standalone/common/function/geo.result b/tests/cases/standalone/common/function/geo.result index 31efbf401da..3fc73a3b98b 100644 --- a/tests/cases/standalone/common/function/geo.result +++ b/tests/cases/standalone/common/function/geo.result @@ -437,3 +437,47 @@ FROM | true | false | true | false | false | true | +--------------------------+--------------------------+------------------------+------------------------+----------------------------------+----------------------------------+ +CREATE TABLE geo_path_partitioned ( + ts TIMESTAMP TIME INDEX, + k INT, + lat DOUBLE, + lon DOUBLE, + PRIMARY KEY(k) +) +PARTITION ON COLUMNS (k) (k < 10, k >= 10 AND k < 20, k >= 20); + +Affected Rows: 0 + +INSERT INTO geo_path_partitioned VALUES + (1000, 1, 1, 11), + (2000, 11, 2, 12), + (3000, 21, 3, 13), + (4000, 2, 4, 14), + (5000, 12, 5, 15), + (6000, 22, 6, 16), + (7000, 3, 7, 17); + +Affected Rows: 7 + +SELECT lat > 3 AS g, geo_path(lat, lon, ts) FROM geo_path_partitioned GROUP BY g ORDER BY g; + ++-------+-------------------------------------------------------------------------------------+ +| g | geo_path(geo_path_partitioned.lat,geo_path_partitioned.lon,geo_path_partitioned.ts) | ++-------+-------------------------------------------------------------------------------------+ +| false | {lat: [1.0, 2.0, 3.0], lng: [11.0, 12.0, 13.0]} | +| true | {lat: [4.0, 5.0, 6.0, 7.0], lng: [14.0, 15.0, 16.0, 17.0]} | ++-------+-------------------------------------------------------------------------------------+ + +SELECT lat > 3 AS g, json_encode_path(lat, lon, ts) FROM geo_path_partitioned GROUP BY g ORDER BY g; + ++-------+---------------------------------------------------------------------------------------------+ +| g | json_encode_path(geo_path_partitioned.lat,geo_path_partitioned.lon,geo_path_partitioned.ts) | ++-------+---------------------------------------------------------------------------------------------+ +| false | [[11.0,1.0],[12.0,2.0],[13.0,3.0]] | +| true | [[14.0,4.0],[15.0,5.0],[16.0,6.0],[17.0,7.0]] | ++-------+---------------------------------------------------------------------------------------------+ + +DROP TABLE geo_path_partitioned; + +Affected Rows: 0 + diff --git a/tests/cases/standalone/common/function/geo.sql b/tests/cases/standalone/common/function/geo.sql index 43025cc196b..1de7d7c82aa 100644 --- a/tests/cases/standalone/common/function/geo.sql +++ b/tests/cases/standalone/common/function/geo.sql @@ -180,3 +180,27 @@ FROM 'POLYGON ((-121.491698 38.653343, -121.582353 38.556757, -121.469721 38.449287, -121.315883 38.541721, -121.491698 38.653343))' AS polygon2, 'POLYGON ((-122.089628 37.450332, -122.20535 37.378342, -122.093062 37.36088, -122.044301 37.372886, -122.089628 37.450332))' AS polygon3, ); + +CREATE TABLE geo_path_partitioned ( + ts TIMESTAMP TIME INDEX, + k INT, + lat DOUBLE, + lon DOUBLE, + PRIMARY KEY(k) +) +PARTITION ON COLUMNS (k) (k < 10, k >= 10 AND k < 20, k >= 20); + +INSERT INTO geo_path_partitioned VALUES + (1000, 1, 1, 11), + (2000, 11, 2, 12), + (3000, 21, 3, 13), + (4000, 2, 4, 14), + (5000, 12, 5, 15), + (6000, 22, 6, 16), + (7000, 3, 7, 17); + +SELECT lat > 3 AS g, geo_path(lat, lon, ts) FROM geo_path_partitioned GROUP BY g ORDER BY g; + +SELECT lat > 3 AS g, json_encode_path(lat, lon, ts) FROM geo_path_partitioned GROUP BY g ORDER BY g; + +DROP TABLE geo_path_partitioned;