use std::sync::Arc;
use serde::{Deserialize, Serialize};
use serde_json::Value;
use tokio::sync::mpsc;
use crate::mapreduce::job::MapReduceJob;
use crate::mapreduce::phase::Phase;
use crate::mapreduce::registry::PhaseRegistry;
#[derive(Debug, thiserror::Error)]
#[non_exhaustive]
pub enum MrError {
#[error("unknown {kind} function: {name}")]
UnknownFunction {
kind: &'static str,
name: String,
},
#[error("phase {phase} ({kind}) failed: {message}")]
PhaseFailed {
phase: u32,
kind: &'static str,
message: String,
},
#[error("unsupported MapReduce inputs: {0}")]
UnsupportedInputs(&'static str),
#[error("wasm-module phases are not enabled on this executor")]
WasmNotImplemented,
#[error("wasm module not found: {0}")]
WasmModuleNotFound(String),
#[error("wasm phase exceeded execution time / fuel limit")]
WasmExecutionTimeout,
#[error("wasm phase exceeded memory limit")]
WasmMemoryLimit,
#[error("wasm runtime error: {0}")]
WasmRuntime(String),
#[error("wasm phase encoding error: {0}")]
WasmEncoding(String),
#[error("link phases are not implemented in this slice")]
LinkNotImplemented,
#[error("internal pipeline error: {0}")]
Pipeline(String),
#[error("json: {0}")]
Json(String),
}
pub trait WasmHook: Send + Sync {
fn apply_phase(
&self,
module_id: &str,
fn_name: &str,
inputs: &[Value],
) -> Result<Vec<Value>, MrError>;
}
#[derive(Clone, Debug, Eq, PartialEq, Serialize, Deserialize)]
pub struct PhaseOutput {
pub phase: u32,
pub value: Value,
}
#[derive(Clone, Debug, Eq, PartialEq, Serialize, Deserialize)]
pub struct PhaseBatch {
pub phase: u32,
pub data: Vec<Value>,
}
pub async fn run_job(
job: MapReduceJob,
registry: Arc<PhaseRegistry>,
) -> Result<Vec<PhaseOutput>, MrError> {
run_job_with_wasm(job, registry, None).await
}
#[must_use]
pub fn run_job_streaming(
job: MapReduceJob,
registry: Arc<PhaseRegistry>,
) -> mpsc::Receiver<Result<PhaseBatch, MrError>> {
run_job_streaming_with_wasm(job, registry, None)
}
#[must_use]
pub fn run_job_streaming_with_wasm(
job: MapReduceJob,
registry: Arc<PhaseRegistry>,
wasm: Option<Arc<dyn WasmHook>>,
) -> mpsc::Receiver<Result<PhaseBatch, MrError>> {
let (tx, rx) = mpsc::channel::<Result<PhaseBatch, MrError>>(4);
tokio::spawn(async move {
let result = stream_job_inner(job, registry, wasm, tx.clone()).await;
if let Err(e) = result {
let _ = tx.send(Err(e)).await;
}
});
rx
}
async fn stream_job_inner(
job: MapReduceJob,
registry: Arc<PhaseRegistry>,
wasm: Option<Arc<dyn WasmHook>>,
tx: mpsc::Sender<Result<PhaseBatch, MrError>>,
) -> Result<(), MrError> {
let items = job
.inputs
.items()
.ok_or(MrError::UnsupportedInputs("bucket scan"))?;
let initial: Vec<Value> = items.into_iter().map(|kd| kd.to_value()).collect();
materialised_initial_inputs_must_be_iterable(&initial);
if job.phases.is_empty() {
if !initial.is_empty() {
let batch = PhaseBatch {
phase: 0,
data: initial,
};
let _ = tx.send(Ok(batch)).await;
}
return Ok(());
}
let n_phases = job.phases.len();
let mut current: Vec<Value> = initial;
for (idx, phase) in job.phases.iter().enumerate() {
let phase_idx = u32::try_from(idx)
.map_err(|_| MrError::Pipeline("phase index exceeds u32 range".into()))?;
let is_last = idx + 1 == n_phases;
let outputs = run_phase(phase_idx, phase, current, ®istry, wasm.as_ref()).await?;
if (phase.keep() || is_last) && !outputs.is_empty() {
let batch = PhaseBatch {
phase: phase_idx,
data: outputs.clone(),
};
if tx.send(Ok(batch)).await.is_err() {
return Ok(());
}
}
current = outputs;
}
Ok(())
}
pub async fn run_job_with_wasm(
job: MapReduceJob,
registry: Arc<PhaseRegistry>,
wasm: Option<Arc<dyn WasmHook>>,
) -> Result<Vec<PhaseOutput>, MrError> {
let items = job
.inputs
.items()
.ok_or(MrError::UnsupportedInputs("bucket scan"))?;
let initial: Vec<Value> = items.into_iter().map(|kd| kd.to_value()).collect();
materialised_initial_inputs_must_be_iterable(&initial);
if job.phases.is_empty() {
let out: Vec<PhaseOutput> = initial
.into_iter()
.map(|v| PhaseOutput { phase: 0, value: v })
.collect();
return Ok(out);
}
let n_phases = job.phases.len();
let mut current: Vec<Value> = initial;
let mut captured: Vec<PhaseOutput> = Vec::new();
for (idx, phase) in job.phases.iter().enumerate() {
let phase_idx = u32::try_from(idx)
.map_err(|_| MrError::Pipeline("phase index exceeds u32 range".into()))?;
let is_last = idx + 1 == n_phases;
let outputs = run_phase(phase_idx, phase, current, ®istry, wasm.as_ref()).await?;
if phase.keep() || is_last {
for v in &outputs {
captured.push(PhaseOutput {
phase: phase_idx,
value: v.clone(),
});
}
}
current = outputs;
}
Ok(captured)
}
async fn run_phase(
phase_idx: u32,
phase: &Phase,
inputs: Vec<Value>,
registry: &Arc<PhaseRegistry>,
wasm: Option<&Arc<dyn WasmHook>>,
) -> Result<Vec<Value>, MrError> {
let (tx_in, rx_in) = mpsc::channel::<Value>(64);
let (tx_out, mut rx_out) = mpsc::channel::<Value>(64);
tokio::spawn(async move {
for v in inputs {
if tx_in.send(v).await.is_err() {
return;
}
}
});
let phase_clone = phase.clone();
let registry_clone = Arc::clone(registry);
let wasm_clone = wasm.cloned();
let phase_join = tokio::spawn(async move {
run_phase_task(
phase_idx,
phase_clone,
rx_in,
tx_out,
registry_clone,
wasm_clone,
)
.await
});
let mut outputs = Vec::new();
while let Some(v) = rx_out.recv().await {
outputs.push(v);
}
match phase_join.await {
Ok(Ok(())) => Ok(outputs),
Ok(Err(e)) => Err(e),
Err(e) => Err(MrError::Pipeline(format!("join: {e}"))),
}
}
async fn run_phase_task(
phase_idx: u32,
phase: Phase,
mut rx: mpsc::Receiver<Value>,
tx: mpsc::Sender<Value>,
registry: Arc<PhaseRegistry>,
wasm: Option<Arc<dyn WasmHook>>,
) -> Result<(), MrError> {
match phase {
Phase::Map { fn_name, arg, .. } => {
let f = registry.map_fn(&fn_name).ok_or(MrError::UnknownFunction {
kind: "map",
name: fn_name.clone(),
})?;
let f = f.clone();
while let Some(v) = rx.recv().await {
let outs = (f)(&v, arg.as_ref()).map_err(|e| MrError::PhaseFailed {
phase: phase_idx,
kind: "map",
message: e.to_string(),
})?;
for o in outs {
if tx.send(o).await.is_err() {
return Err(MrError::Pipeline(
"downstream phase dropped its inbound channel".into(),
));
}
}
}
Ok(())
}
Phase::Reduce { fn_name, arg, .. } => {
let f = registry
.reduce_fn(&fn_name)
.ok_or(MrError::UnknownFunction {
kind: "reduce",
name: fn_name.clone(),
})?;
let f = f.clone();
let mut buf: Vec<Value> = Vec::new();
while let Some(v) = rx.recv().await {
buf.push(v);
}
let outs = (f)(&buf, arg.as_ref()).map_err(|e| MrError::PhaseFailed {
phase: phase_idx,
kind: "reduce",
message: e.to_string(),
})?;
for o in outs {
if tx.send(o).await.is_err() {
return Err(MrError::Pipeline(
"downstream phase dropped its inbound channel".into(),
));
}
}
Ok(())
}
Phase::Link { .. } => {
Err(MrError::LinkNotImplemented)
}
Phase::WasmModule {
module_id, fn_name, ..
} => {
let hook = wasm.ok_or(MrError::WasmNotImplemented)?;
let mut buf: Vec<Value> = Vec::new();
while let Some(v) = rx.recv().await {
buf.push(v);
}
let mid = module_id.clone();
let fname = fn_name.clone();
let outs = tokio::task::spawn_blocking(move || hook.apply_phase(&mid, &fname, &buf))
.await
.map_err(|e| MrError::Pipeline(format!("wasm join: {e}")))?
.map_err(|e| match e {
MrError::WasmModuleNotFound(_)
| MrError::WasmExecutionTimeout
| MrError::WasmMemoryLimit
| MrError::WasmRuntime(_)
| MrError::WasmEncoding(_)
| MrError::WasmNotImplemented => e,
other => MrError::PhaseFailed {
phase: phase_idx,
kind: "wasm",
message: other.to_string(),
},
})?;
for o in outs {
if tx.send(o).await.is_err() {
return Err(MrError::Pipeline(
"downstream phase dropped its inbound channel".into(),
));
}
}
Ok(())
}
}
}
fn materialised_initial_inputs_must_be_iterable<T>(_: &[T]) {}
#[cfg(test)]
mod tests {
use super::*;
use crate::mapreduce::builtins::default_registry;
use crate::mapreduce::job::{Inputs, KeyDatum};
fn registry() -> Arc<PhaseRegistry> {
Arc::new(default_registry())
}
#[tokio::test]
async fn empty_phase_list_is_identity() {
let job = MapReduceJob {
inputs: Inputs::KeyData(vec![KeyDatum::with_value("b", "k", serde_json::json!(1))]),
phases: vec![],
timeout_ms: None,
};
let out = run_job(job, registry()).await.expect("ok");
assert_eq!(out.len(), 1);
assert_eq!(out[0].phase, 0);
assert_eq!(out[0].value["bucket"], "b");
}
#[tokio::test]
async fn map_then_reduce_pipeline() {
let job = MapReduceJob {
inputs: Inputs::KeyData(vec![
KeyDatum::with_value("b", "k1", serde_json::json!(2)),
KeyDatum::with_value("b", "k2", serde_json::json!(3)),
KeyDatum::with_value("b", "k3", serde_json::json!(4)),
]),
phases: vec![
Phase::Map {
fn_name: "map_object_value".into(),
arg: None,
keep: false,
},
Phase::Reduce {
fn_name: "reduce_sum".into(),
arg: None,
keep: true,
},
],
timeout_ms: None,
};
let out = run_job(job, registry()).await.expect("ok");
assert_eq!(out.len(), 1);
assert_eq!(out[0].phase, 1);
assert_eq!(out[0].value, serde_json::json!(9));
}
#[tokio::test]
async fn keep_intermediate_phase_outputs_are_captured() {
let job = MapReduceJob {
inputs: Inputs::KeyData(vec![
KeyDatum::with_value("b", "k1", serde_json::json!(5)),
KeyDatum::with_value("b", "k2", serde_json::json!(7)),
]),
phases: vec![
Phase::Map {
fn_name: "map_object_value".into(),
arg: None,
keep: true,
},
Phase::Reduce {
fn_name: "reduce_sum".into(),
arg: None,
keep: true,
},
],
timeout_ms: None,
};
let out = run_job(job, registry()).await.expect("ok");
assert_eq!(out.len(), 3);
assert_eq!(out[0].phase, 0);
assert_eq!(out[1].phase, 0);
assert_eq!(out[2].phase, 1);
assert_eq!(out[2].value, serde_json::json!(12));
}
#[tokio::test]
async fn unknown_map_function_is_typed_error() {
let job = MapReduceJob {
inputs: Inputs::KeyData(vec![KeyDatum::with_value("b", "k", serde_json::json!(1))]),
phases: vec![Phase::Map {
fn_name: "no_such_function".into(),
arg: None,
keep: false,
}],
timeout_ms: None,
};
let err = run_job(job, registry()).await.expect_err("error");
assert!(matches!(err, MrError::UnknownFunction { kind: "map", .. }));
}
#[tokio::test]
async fn unknown_reduce_function_is_typed_error() {
let job = MapReduceJob {
inputs: Inputs::KeyData(vec![]),
phases: vec![Phase::Reduce {
fn_name: "no_such_reduce".into(),
arg: None,
keep: false,
}],
timeout_ms: None,
};
let err = run_job(job, registry()).await.expect_err("error");
assert!(matches!(
err,
MrError::UnknownFunction { kind: "reduce", .. }
));
}
#[tokio::test]
async fn bucket_inputs_unsupported() {
let job = MapReduceJob {
inputs: Inputs::Bucket("b".into()),
phases: vec![],
timeout_ms: None,
};
let err = run_job(job, registry()).await.expect_err("error");
assert!(matches!(err, MrError::UnsupportedInputs(_)));
}
#[tokio::test]
async fn wasm_phase_returns_typed_error() {
let job = MapReduceJob {
inputs: Inputs::KeyData(vec![KeyDatum::with_value("b", "k", serde_json::json!(1))]),
phases: vec![Phase::WasmModule {
module_id: "m".into(),
fn_name: "f".into(),
arg: None,
keep: false,
}],
timeout_ms: None,
};
let err = run_job(job, registry()).await.expect_err("error");
assert!(matches!(err, MrError::WasmNotImplemented));
}
#[tokio::test]
async fn link_phase_returns_typed_error() {
let job = MapReduceJob {
inputs: Inputs::KeyData(vec![KeyDatum::pair("b", "k")]),
phases: vec![Phase::Link {
bucket: None,
tag: None,
keep: false,
}],
timeout_ms: None,
};
let err = run_job(job, registry()).await.expect_err("error");
assert!(matches!(err, MrError::LinkNotImplemented));
}
#[tokio::test]
async fn determinism_under_one_hundred_inputs_three_phases() {
let mut data = Vec::new();
for i in 0..100u32 {
data.push(KeyDatum::with_value(
"b",
format!("k{i}"),
serde_json::json!(i),
));
}
let phases = vec![
Phase::Map {
fn_name: "map_object_value".into(),
arg: None,
keep: false,
},
Phase::Reduce {
fn_name: "reduce_sort".into(),
arg: None,
keep: false,
},
Phase::Reduce {
fn_name: "reduce_count".into(),
arg: None,
keep: true,
},
];
let job = MapReduceJob {
inputs: Inputs::KeyData(data.clone()),
phases: phases.clone(),
timeout_ms: None,
};
let job2 = MapReduceJob {
inputs: Inputs::KeyData(data),
phases,
timeout_ms: None,
};
let r1 = run_job(job, registry()).await.expect("run1");
let r2 = run_job(job2, registry()).await.expect("run2");
assert_eq!(r1, r2);
assert_eq!(r1.len(), 1);
assert_eq!(r1[0].value, serde_json::json!(100));
}
#[tokio::test]
async fn pairs_inputs_are_materialised_with_null_value() {
let job = MapReduceJob {
inputs: Inputs::Pairs(vec![("b".into(), "k1".into()), ("b".into(), "k2".into())]),
phases: vec![],
timeout_ms: None,
};
let out = run_job(job, registry()).await.expect("ok");
assert_eq!(out.len(), 2);
assert!(out[0].value["value"].is_null());
assert_eq!(out[0].value["bucket"], "b");
assert_eq!(out[0].value["key"], "k1");
}
async fn drain_stream(
mut rx: mpsc::Receiver<Result<PhaseBatch, MrError>>,
) -> Vec<Result<PhaseBatch, MrError>> {
let mut out = Vec::new();
while let Some(item) = rx.recv().await {
out.push(item);
}
out
}
#[tokio::test]
async fn streaming_emits_one_batch_per_kept_phase() {
let job = MapReduceJob {
inputs: Inputs::KeyData(vec![
KeyDatum::with_value("b", "k1", serde_json::json!(1)),
KeyDatum::with_value("b", "k2", serde_json::json!(2)),
KeyDatum::with_value("b", "k3", serde_json::json!(3)),
]),
phases: vec![
Phase::Map {
fn_name: "map_object_value".into(),
arg: None,
keep: true,
},
Phase::Reduce {
fn_name: "reduce_sum".into(),
arg: None,
keep: true,
},
],
timeout_ms: None,
};
let rx = run_job_streaming(job, registry());
let items = drain_stream(rx).await;
assert_eq!(items.len(), 2, "two kept phases yield two batches");
let b0 = items[0].as_ref().expect("phase 0 ok");
assert_eq!(b0.phase, 0);
assert_eq!(b0.data.len(), 3);
let b1 = items[1].as_ref().expect("phase 1 ok");
assert_eq!(b1.phase, 1);
assert_eq!(b1.data, vec![serde_json::json!(6)]);
}
#[tokio::test]
async fn streaming_skips_non_keep_intermediate_phases() {
let job = MapReduceJob {
inputs: Inputs::KeyData(vec![KeyDatum::with_value("b", "k", serde_json::json!(2))]),
phases: vec![
Phase::Map {
fn_name: "map_object_value".into(),
arg: None,
keep: false,
},
Phase::Reduce {
fn_name: "reduce_sum".into(),
arg: None,
keep: true,
},
],
timeout_ms: None,
};
let rx = run_job_streaming(job, registry());
let items = drain_stream(rx).await;
assert_eq!(items.len(), 1, "only the kept reduce emits a batch");
let b = items[0].as_ref().expect("ok");
assert_eq!(b.phase, 1);
}
#[tokio::test]
async fn streaming_surfaces_phase_error_as_terminal_item() {
let job = MapReduceJob {
inputs: Inputs::KeyData(vec![KeyDatum::pair("b", "k")]),
phases: vec![Phase::Map {
fn_name: "no_such_function".into(),
arg: None,
keep: true,
}],
timeout_ms: None,
};
let rx = run_job_streaming(job, registry());
let items = drain_stream(rx).await;
assert_eq!(items.len(), 1);
let err = items[0].as_ref().expect_err("unknown function");
assert!(matches!(err, MrError::UnknownFunction { kind: "map", .. }));
}
#[tokio::test]
async fn streaming_empty_phase_list_emits_initial_inputs() {
let job = MapReduceJob {
inputs: Inputs::KeyData(vec![KeyDatum::with_value("b", "k", serde_json::json!(42))]),
phases: vec![],
timeout_ms: None,
};
let rx = run_job_streaming(job, registry());
let items = drain_stream(rx).await;
assert_eq!(items.len(), 1);
let b = items[0].as_ref().expect("ok");
assert_eq!(b.phase, 0);
assert_eq!(b.data.len(), 1);
assert_eq!(b.data[0]["value"], serde_json::json!(42));
}
}