use std::collections::HashMap;
use std::path::Path;
use super::state_emit::AGG_STATE_FIELD;
use crate::bridge::envelope::{ErrorCode, Response};
use crate::data::executor::core_loop::CoreLoop;
use crate::data::executor::handlers::accum::GroupState;
use crate::data::executor::handlers::join::FrameStreamReader;
use crate::data::executor::task::ExecutionTask;
use nodedb_physical::physical_plan::{AggregateSpec, GroupKeySpec};
use nodedb_query::msgpack_scan;
fn row_key_specs(group_by: &[GroupKeySpec]) -> Vec<GroupKeySpec> {
group_by
.iter()
.map(|s| GroupKeySpec::column(s.output_name.clone()))
.collect()
}
pub(in crate::data::executor) fn merge_state_frames(
state_path: &Path,
group_by: &[GroupKeySpec],
_aggregates: &[AggregateSpec],
) -> crate::Result<HashMap<String, GroupState>> {
let mut merged: HashMap<String, GroupState> = HashMap::new();
let mut reader = FrameStreamReader::open(state_path)?;
let row_keys = row_key_specs(group_by);
while let Some(row) = reader.next_row()? {
let key = msgpack_scan::build_group_key(&row, &row_keys);
let (val_start, _val_end) = msgpack_scan::extract_field(&row, 0, AGG_STATE_FIELD)
.ok_or_else(|| crate::Error::Codec {
detail: format!("shuffle-aggregate row missing `{AGG_STATE_FIELD}` field"),
})?;
let mut off = val_start;
let state_bytes =
msgpack_scan::read_bin_advance(&row, &mut off).ok_or_else(|| crate::Error::Codec {
detail: format!("shuffle-aggregate `{AGG_STATE_FIELD}` is not binary"),
})?;
let state: GroupState =
sonic_rs::from_slice(state_bytes).map_err(|e| crate::Error::Codec {
detail: format!("shuffle-aggregate partial-state decode: {e}"),
})?;
match merged.entry(key) {
std::collections::hash_map::Entry::Occupied(mut o) => {
o.get_mut().merge_from(state);
}
std::collections::hash_map::Entry::Vacant(v) => {
v.insert(state);
}
}
}
Ok(merged)
}
pub(in crate::data::executor) struct ShuffleAggregateParams<'a> {
pub task: &'a ExecutionTask,
pub state_path: &'a str,
pub group_by: &'a [GroupKeySpec],
pub aggregates: &'a [AggregateSpec],
pub having: &'a [u8],
pub limit: usize,
pub sort_keys: &'a [(String, bool)],
}
impl CoreLoop {
pub(in crate::data::executor) fn execute_shuffle_aggregate(
&mut self,
params: ShuffleAggregateParams<'_>,
) -> Response {
let ShuffleAggregateParams {
task,
state_path,
group_by,
aggregates,
having,
limit,
sort_keys,
} = params;
let merged = match merge_state_frames(Path::new(state_path), group_by, aggregates) {
Ok(m) => m,
Err(e) => {
return self.response_error(
task,
ErrorCode::Internal {
detail: e.to_string(),
},
);
}
};
match self.finalize_groups(super::streaming::finalize::FinalizeGroupsParams {
groups: merged,
sub_groups: HashMap::new(),
group_by,
aggregates,
having,
limit,
sub_group_by: &[],
sub_aggregates: &[],
sort_keys,
}) {
Ok(payload) => self.response_with_payload(task, payload),
Err(e) => self.response_error(
task,
ErrorCode::Internal {
detail: e.to_string(),
},
),
}
}
}
#[cfg(test)]
mod tests {
use std::collections::HashMap;
use std::io::Write as _;
use super::merge_state_frames;
use crate::data::executor::handlers::accum::GroupState;
use nodedb_physical::physical_plan::{AggregateSpec, GroupKeySpec};
use nodedb_types::Value;
fn make_spec(func: &str, field: &str) -> AggregateSpec {
AggregateSpec {
function: func.to_string(),
field: field.to_string(),
alias: format!("{func}({field})"),
user_alias: None,
expr: None,
}
}
fn make_doc(group: &str, value: Option<i64>) -> Vec<u8> {
let mut map: HashMap<String, Value> = HashMap::new();
map.insert("g".to_string(), Value::String(group.to_string()));
if let Some(v) = value {
map.insert("v".to_string(), Value::Integer(v));
}
nodedb_types::value_to_msgpack(&Value::Object(map)).expect("encode doc")
}
fn write_partial_state_frames(
path: &std::path::Path,
specs: &[AggregateSpec],
group_by: &[GroupKeySpec],
docs: &[Vec<u8>],
) {
let mut groups: HashMap<String, GroupState> = HashMap::new();
for doc in docs {
let key = nodedb_query::msgpack_scan::build_group_key(doc, group_by);
groups
.entry(key)
.or_insert_with(|| GroupState::new(specs))
.feed(specs, doc);
}
let mut f = std::fs::File::create(path).expect("create frame file");
for (key, state) in groups {
let parts: Vec<serde_json::Value> = sonic_rs::from_str(&key).expect("key json");
let mut row: HashMap<String, Value> = HashMap::new();
let mut part_idx = 0usize;
for spec in group_by {
if spec.field.is_none() {
continue;
}
let jv = parts
.get(part_idx)
.cloned()
.unwrap_or(serde_json::Value::Null);
row.insert(spec.output_name.clone(), Value::from(jv));
part_idx += 1;
}
let state_bytes = sonic_rs::to_vec(&state).expect("state json");
row.insert(
super::AGG_STATE_FIELD.to_string(),
Value::Bytes(state_bytes),
);
let row_bytes =
nodedb_types::value_to_msgpack(&Value::Object(row)).expect("encode row");
let len = u32::try_from(row_bytes.len()).expect("row fits u32");
f.write_all(&len.to_le_bytes()).expect("write len");
f.write_all(&row_bytes).expect("write body");
}
f.flush().expect("flush");
}
fn reference_finalized(
specs: &[AggregateSpec],
group_by: &[GroupKeySpec],
docs: &[Vec<u8>],
) -> HashMap<String, Vec<Value>> {
let mut map: HashMap<String, GroupState> = HashMap::new();
for doc in docs {
let key = nodedb_query::msgpack_scan::build_group_key(doc, group_by);
map.entry(key)
.or_insert_with(|| GroupState::new(specs))
.feed(specs, doc);
}
map.into_iter()
.map(|(k, s)| (k, s.finalize(specs).into_iter().map(|(_, v)| v).collect()))
.collect()
}
#[test]
fn shuffle_combiner_equals_single_pass_over_union() {
let specs = vec![
make_spec("count", "*"),
make_spec("count", "v"),
make_spec("sum", "v"),
make_spec("avg", "v"),
make_spec("min", "v"),
make_spec("max", "v"),
make_spec("stddev_pop", "v"),
make_spec("count_distinct", "v"),
];
let group_by = vec![GroupKeySpec::column("g")];
let docs_a: Vec<Vec<u8>> = vec![
make_doc("x", Some(1)),
make_doc("x", Some(2)),
make_doc("y", Some(10)),
make_doc("x", None),
make_doc("y", Some(10)), ];
let docs_b: Vec<Vec<u8>> = vec![
make_doc("x", Some(3)),
make_doc("y", Some(20)),
make_doc("y", Some(30)),
make_doc("z", Some(7)),
make_doc("z", Some(7)),
];
let dir = tempfile::tempdir().expect("tempdir");
let path_a = dir.path().join("producer_a.frames");
let path_b = dir.path().join("producer_b.frames");
write_partial_state_frames(&path_a, &specs, &group_by, &docs_a);
write_partial_state_frames(&path_b, &specs, &group_by, &docs_b);
let combined = dir.path().join("combined.frames");
{
let mut out = std::fs::File::create(&combined).expect("create combined");
for p in [&path_a, &path_b] {
let bytes = std::fs::read(p).expect("read producer frames");
out.write_all(&bytes).expect("append");
}
out.flush().expect("flush combined");
}
let merged = merge_state_frames(&combined, &group_by, &specs).expect("merge");
let got: HashMap<String, Vec<Value>> = merged
.into_iter()
.map(|(k, s)| (k, s.finalize(&specs).into_iter().map(|(_, v)| v).collect()))
.collect();
let mut union = docs_a.clone();
union.extend(docs_b.clone());
let expected = reference_finalized(&specs, &group_by, &union);
assert_eq!(
got.len(),
expected.len(),
"group count must match the single-pass reference"
);
for (k, ev) in &expected {
let gv = got
.get(k)
.unwrap_or_else(|| panic!("merged result missing group {k}"));
assert_eq!(
gv, ev,
"aggregate values for group {k} must equal the single-pass union aggregate"
);
}
}
}