fix: merge every state row in geo_path and json_encode_path (#9339)

Signed-off-by: Dennis Zhuang <killme2008@gmail.com>
This commit is contained in:
dennis zhuang
2026-09-24 02:12:11 +00:00
committed by GitHub
parent c7fa48ef95
commit cc1be82858
4 changed files with 108 additions and 28 deletions
+20 -14
View File
@@ -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::<Float64Type>(&lat_list)?;
let lng_list = as_list_array(state.column(1))?.value(0);
let lng_array = as_primitive_array::<Float64Type>(&lng_list)?;
let ts_list = as_list_array(state.column(2))?.value(0);
let ts_array = as_primitive_array::<Int64Type>(&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::<Float64Type>(&lat_list)?);
self.lng
.extend(as_primitive_array::<Float64Type>(&lng_list)?);
self.timestamp
.extend(as_primitive_array::<Int64Type>(&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();
+20 -14
View File
@@ -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::<Float64Type>(&lat_list)?;
let lng_list = as_list_array(state.column(1))?.value(0);
let lng_array = as_primitive_array::<Float64Type>(&lng_list)?;
let ts_list = as_list_array(state.column(2))?.value(0);
let ts_array = as_primitive_array::<Int64Type>(&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::<Float64Type>(&lat_list)?);
self.lng
.extend(as_primitive_array::<Float64Type>(&lng_list)?);
self.timestamp
.extend(as_primitive_array::<Int64Type>(&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();
@@ -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
@@ -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;